已合并
engram support training #8914
engram support training #8914
已合并
luozhonglin创建于 7月20日
共 23 个文件变更+2508-351
@@ -24,7 +24,9 @@
24 24 
25using namespace op;25using namespace op;
26 26 
27-static bool CheckNotNull(const aclTensor *commContext, const aclTensor *indices, aclTensor *fetched)27+static bool CheckNotNull(const aclTensor *commContext, const aclTensor *indices, const aclTensor *localStorageAddr,
28+ aclTensor *fetched, aclTensor *permOut, aclTensor *sendCountsOut, aclTensor *recvCountsOut,
29+ aclTensor *recvLocalEntryOut, aclTensor *numRecvOut)
28{30{
29 OP_CHECK_NULL(commContext, return false);31 OP_CHECK_NULL(commContext, return false);
30 OP_CHECK_NULL(indices, return false);32 OP_CHECK_NULL(indices, return false);
@@ -32,9 +34,14 @@ static bool CheckNotNull(const aclTensor *commContext, const aclTensor *indices,
32 return true;34 return true;
33}35}
34 36 
35-static aclnnStatus CheckParams(const aclTensor *commContext, const aclTensor *indices, aclTensor *fetched)37+static aclnnStatus CheckParams(const aclTensor *commContext, const aclTensor *indices,
38+ const aclTensor *localStorageAddr, aclTensor *fetched, aclTensor *permOut,
39+ aclTensor *sendCountsOut, aclTensor *recvCountsOut, aclTensor *recvLocalEntryOut,
40+ aclTensor *numRecvOut)
36{41{
37- CHECK_RET(CheckNotNull(commContext, indices, fetched), ACLNN_ERR_PARAM_NULLPTR);42+ CHECK_RET(CheckNotNull(commContext, indices, localStorageAddr, fetched, permOut, sendCountsOut, recvCountsOut,
43+ recvLocalEntryOut, numRecvOut),
44+ ACLNN_ERR_PARAM_NULLPTR);
38 return ACLNN_SUCCESS;45 return ACLNN_SUCCESS;
39}46}
40 47 
@@ -42,14 +49,21 @@ static aclnnStatus CheckParams(const aclTensor *commContext, const aclTensor *in
42extern "C" {49extern "C" {
43#endif50#endif
44 51 
45-aclnnStatus aclnnEngramFetchGetWorkspaceSize(const aclTensor *commContext, const aclTensor *indices, int32_t hiddenSize,52+aclnnStatus aclnnEngramFetchGetWorkspaceSize(const aclTensor *commContext, const aclTensor *indices,
46- int64_t numEntriesPerRank, aclTensor *fetched, uint64_t *workspaceSize,53+ const aclTensor *localStorageAddr, aclTensor *fetched, aclTensor *permOut,
54+ aclTensor *sendCountsOut, aclTensor *recvCountsOut,
55+ aclTensor *recvLocalEntryOut, aclTensor *numRecvOut, int32_t hiddenSize,
56+ int64_t numEntriesPerRank, int64_t numMaxTokensPerRank,
57+ int64_t commBufferSize, int64_t withGrad, uint64_t *workspaceSize,
47 aclOpExecutor **executor)58 aclOpExecutor **executor)
48{59{
49- auto retParam = CheckParams(commContext, indices, fetched);60+ auto retParam = CheckParams(commContext, indices, localStorageAddr, fetched, permOut, sendCountsOut, recvCountsOut,
61+ recvLocalEntryOut, numRecvOut);
50 CHECK_RET(retParam == ACLNN_SUCCESS, retParam);62 CHECK_RET(retParam == ACLNN_SUCCESS, retParam);
51- aclnnStatus ret = aclnnInnerEngramFetchGetWorkspaceSize(commContext, indices, hiddenSize, numEntriesPerRank,63+ aclnnStatus ret = aclnnInnerEngramFetchGetWorkspaceSize(commContext, indices, localStorageAddr, hiddenSize,
52- fetched, workspaceSize, executor);64+ numEntriesPerRank, numMaxTokensPerRank, commBufferSize,
65+ withGrad, fetched, permOut, sendCountsOut, recvCountsOut,
66+ recvLocalEntryOut, numRecvOut, workspaceSize, executor);
53 return ret;67 return ret;
54}68}
55 69 
@@ -20,17 +20,27 @@ class EngramFetch : public OpDef {
20public:20public:
21 explicit EngramFetch(const char *name) : OpDef(name)21 explicit EngramFetch(const char *name) : OpDef(name)
22 {22 {
23- this->Input("commContext").ParamType(REQUIRED).DataTypeList({ge::DT_INT32}).FormatList({ge::FORMAT_ND});23+ this->Input("comm_context").ParamType(REQUIRED).DataTypeList({ge::DT_INT32}).FormatList({ge::FORMAT_ND});
24- 
25 this->Input("indices").ParamType(REQUIRED).DataTypeList({ge::DT_INT32}).FormatList({ge::FORMAT_ND});24 this->Input("indices").ParamType(REQUIRED).DataTypeList({ge::DT_INT32}).FormatList({ge::FORMAT_ND});
25+ this->Input("local_storage_addr").ParamType(OPTIONAL).DataTypeList({ge::DT_INT64}).FormatList({ge::FORMAT_ND});
26 26 
27 this->Output("fetched")27 this->Output("fetched")
28 .ParamType(REQUIRED)28 .ParamType(REQUIRED)
29 .DataTypeList({ge::DT_BF16, ge::DT_FLOAT16, ge::DT_FLOAT})29 .DataTypeList({ge::DT_BF16, ge::DT_FLOAT16, ge::DT_FLOAT})
30 .FormatList({ge::FORMAT_ND});30 .FormatList({ge::FORMAT_ND});
31- 31+ this->Output("perm_out").ParamType(OPTIONAL).DataTypeList({ge::DT_INT32}).FormatList({ge::FORMAT_ND});
32+ this->Output("send_counts_out").ParamType(OPTIONAL).DataTypeList({ge::DT_INT32}).FormatList({ge::FORMAT_ND});
33+ this->Output("recv_counts_out").ParamType(OPTIONAL).DataTypeList({ge::DT_INT32}).FormatList({ge::FORMAT_ND});
34+ this->Output("recv_local_entry_out")
35+ .ParamType(OPTIONAL)
36+ .DataTypeList({ge::DT_INT32})
37+ .FormatList({ge::FORMAT_ND});
38+ this->Output("num_recv_out").ParamType(OPTIONAL).DataTypeList({ge::DT_INT32}).FormatList({ge::FORMAT_ND});
32 this->Attr("hidden_size").AttrType(REQUIRED).Int();39 this->Attr("hidden_size").AttrType(REQUIRED).Int();
33 this->Attr("num_entries_per_rank").AttrType(REQUIRED).Int();40 this->Attr("num_entries_per_rank").AttrType(REQUIRED).Int();
41+ this->Attr("num_max_tokens_per_rank").AttrType(OPTIONAL).Int();
42+ this->Attr("comm_buffer_size").AttrType(OPTIONAL).Int();
43+ this->Attr("with_grad").AttrType(OPTIONAL).Int();
34 44 
35 OpAICoreConfig aicore_config_950;45 OpAICoreConfig aicore_config_950;
36 aicore_config_950.DynamicCompileStaticFlag(true)46 aicore_config_950.DynamicCompileStaticFlag(true)
@@ -26,19 +26,41 @@
26using namespace AscendC;26using namespace AscendC;
27using namespace ge;27using namespace ge;
28 28 
29-namespace MC2Tiling {29+namespace Mc2Tiling {
30constexpr uint32_t COMM_CONTEXT_INDEX = 0U;30constexpr uint32_t COMM_CONTEXT_INDEX = 0U;
31constexpr uint32_t INDICES_INDEX = 1U;31constexpr uint32_t INDICES_INDEX = 1U;
32+constexpr uint32_t LOCAL_STORAGE_ADDR_INDEX = 2U;
32constexpr uint32_t FETCHED_INDEX = 0U;33constexpr uint32_t FETCHED_INDEX = 0U;
34+constexpr uint32_t PERM_OUT_INDEX = 1U;
35+constexpr uint32_t SEND_COUNTS_OUT_INDEX = 2U;
36+constexpr uint32_t RECV_COUNTS_OUT_INDEX = 3U;
37+constexpr uint32_t RECV_LOCAL_ENTRY_OUT_INDEX = 4U;
38+constexpr uint32_t NUM_RECV_OUT_INDEX = 5U;
33 39 
34constexpr uint32_t ATTR_HIDDEN_SIZE_INDEX = 0U;40constexpr uint32_t ATTR_HIDDEN_SIZE_INDEX = 0U;
35constexpr uint32_t ATTR_NUM_ENTRIES_PER_RANK_INDEX = 1U;41constexpr uint32_t ATTR_NUM_ENTRIES_PER_RANK_INDEX = 1U;
42+constexpr uint32_t ATTR_NUM_MAX_TOKENS_PER_RANK_INDEX = 2U;
43+constexpr uint32_t ATTR_COMM_BUFFER_SIZE_INDEX = 3U;
44+constexpr uint32_t ATTR_WITH_GRAD_INDEX = 4U;
36 45 
37constexpr uint32_t DIM_ONE = 1U;46constexpr uint32_t DIM_ONE = 1U;
38constexpr uint32_t DIM_TWO = 2U;47constexpr uint32_t DIM_TWO = 2U;
39constexpr uint32_t SYSTEM_NEED_WORKSPACE = 16U * 1024 * 1024;48constexpr uint32_t SYSTEM_NEED_WORKSPACE = 16U * 1024 * 1024;
40 49 
41constexpr int32_t HIDDEN_SIZE_ALIGN = 128;50constexpr int32_t HIDDEN_SIZE_ALIGN = 128;
51+constexpr int64_t UB_ALIGN = 32;
52+constexpr int64_t FLAG_SCRATCH_SIZE = 32;
53+constexpr int64_t WORKSPACE_ALIGN_2MB = 2 * 1024 * 1024;
54+ 
55+static int64_t CeilDiv(int64_t x, int64_t y)
56+{
57+ return (x + y - 1) / y;
58+}
59+ 
60+static int64_t AlignTo(int64_t x, int64_t y)
61+{
62+ return CeilDiv(x, y) * y;
63+}
42 64 
43static const std::vector<ge::DataType> OUTPUT_DTYPE_LIST = {ge::DT_BF16, ge::DT_FLOAT16, ge::DT_FLOAT};65static const std::vector<ge::DataType> OUTPUT_DTYPE_LIST = {ge::DT_BF16, ge::DT_FLOAT16, ge::DT_FLOAT};
44 66 
@@ -60,6 +82,10 @@ static void PrintEngramFetchTilingData(const EngramFetchTilingData *tilingData,
60 OP_LOGD(nodeName, "hiddenDim is %ld", tilingData->hiddenDim);82 OP_LOGD(nodeName, "hiddenDim is %ld", tilingData->hiddenDim);
61 OP_LOGD(nodeName, "hiddenBytes is %ld", tilingData->hiddenBytes);83 OP_LOGD(nodeName, "hiddenBytes is %ld", tilingData->hiddenBytes);
62 OP_LOGD(nodeName, "aivNum is %u", tilingData->aivNum);84 OP_LOGD(nodeName, "aivNum is %u", tilingData->aivNum);
85+ OP_LOGD(nodeName, "rankSize is %u", tilingData->rankSize);
86+ OP_LOGD(nodeName, "numMaxTokensPerRank is %ld", tilingData->numMaxTokensPerRank);
87+ OP_LOGD(nodeName, "totalRecv is %ld", tilingData->totalRecv);
88+ OP_LOGD(nodeName, "commBufferSize is %ld", tilingData->commBufferSize);
63}89}
64 90 
65/**91/**
@@ -107,6 +133,17 @@ static ge::graphStatus CheckTensorDataType(const gert::TilingContext *context)
107 "The dtype of fetched must be DT_BF16, DT_FLOAT16 or DT_FLOAT."),133 "The dtype of fetched must be DT_BF16, DT_FLOAT16 or DT_FLOAT."),
108 return ge::GRAPH_FAILED);134 return ge::GRAPH_FAILED);
109 135 
136+ auto localStorageAddrDesc = context->GetInputDesc(LOCAL_STORAGE_ADDR_INDEX);
137+ const gert::StorageShape *localStorageAddrShape = context->GetInputShape(LOCAL_STORAGE_ADDR_INDEX);
138+ if (localStorageAddrDesc != nullptr && localStorageAddrShape != nullptr) {
139+ OP_TILING_CHECK(
140+ localStorageAddrDesc->GetDataType() != ge::DT_INT64,
141+ OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(nodeName, "localStorageAddr",
142+ Ops::Base::ToString(localStorageAddrDesc->GetDataType()).c_str(),
143+ "The dtype of localStorageAddr must be DT_INT64."),
144+ return ge::GRAPH_FAILED);
145+ }
146+ 
110 return ge::GRAPH_SUCCESS;147 return ge::GRAPH_SUCCESS;
111}148}
112 149 
@@ -167,6 +204,23 @@ static ge::graphStatus CheckTensorDim(const gert::TilingContext *context, int64_
167 (std::string("dim0 must equal numTokens=") + std::to_string(numTokens)).c_str()),204 (std::string("dim0 must equal numTokens=") + std::to_string(numTokens)).c_str()),
168 return ge::GRAPH_FAILED);205 return ge::GRAPH_FAILED);
169 206 
207+ const gert::StorageShape *localStorageAddrShape = context->GetInputShape(LOCAL_STORAGE_ADDR_INDEX);
208+ if (localStorageAddrShape != nullptr) {
209+ OP_TILING_CHECK(localStorageAddrShape->GetStorageShape().GetDimNum() != DIM_ONE,
210+ OP_LOGE_FOR_INVALID_SHAPEDIM_WITH_REASON(
211+ nodeName, "localStorageAddr",
212+ (std::to_string(localStorageAddrShape->GetStorageShape().GetDimNum()) + "D").c_str(),
213+ "The shape dim of localStorageAddr must be 1D."),
214+ return ge::GRAPH_FAILED);
215+ OP_TILING_CHECK(
216+ localStorageAddrShape->GetStorageShape().GetDim(0) != 1,
217+ OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON(
218+ nodeName, "localStorageAddr",
219+ (std::string("dim0=") + std::to_string(localStorageAddrShape->GetStorageShape().GetDim(0))).c_str(),
220+ "dim0 must be 1."),
221+ return ge::GRAPH_FAILED);
222+ }
223+ 
170 return ge::GRAPH_SUCCESS;224 return ge::GRAPH_SUCCESS;
171}225}
172 226 
@@ -284,6 +338,111 @@ static ge::graphStatus CheckAttrParams(const gert::TilingContext *context)
284 return ge::GRAPH_SUCCESS;338 return ge::GRAPH_SUCCESS;
285}339}
286 340 
341+/**
342+ * @brief 校验训练场景的额外属性和输出shape
343+ */
344+static ge::graphStatus CheckTrainingParams(const gert::TilingContext *context, int64_t numTokens)
345+{
346+ const char *nodeName = context->GetNodeName();
347+ auto attrs = context->GetAttrs();
348+ 
349+ auto numMaxTokensPerRankPtr = attrs->GetAttrPointer<int64_t>(ATTR_NUM_MAX_TOKENS_PER_RANK_INDEX);
350+ OP_TILING_CHECK(numMaxTokensPerRankPtr == nullptr, OP_LOGE_WITH_INVALID_INPUT(nodeName, "num_max_tokens_per_rank"),
351+ return ge::GRAPH_FAILED);
352+ OP_TILING_CHECK(*numMaxTokensPerRankPtr <= 0,
353+ OP_LOGE_FOR_INVALID_VALUE(nodeName, "num_max_tokens_per_rank",
354+ std::to_string(*numMaxTokensPerRankPtr).c_str(), "> 0"),
355+ return ge::GRAPH_FAILED);
356+ OP_TILING_CHECK(*numMaxTokensPerRankPtr < numTokens,
357+ OP_LOGE_FOR_INVALID_VALUE(nodeName, "num_max_tokens_per_rank",
358+ std::to_string(*numMaxTokensPerRankPtr).c_str(),
359+ (std::string(">= numTokens(") + std::to_string(numTokens) + ")").c_str()),
360+ return ge::GRAPH_FAILED);
361+ 
362+ auto commBufferSizePtr = attrs->GetAttrPointer<int64_t>(ATTR_COMM_BUFFER_SIZE_INDEX);
363+ OP_TILING_CHECK(commBufferSizePtr == nullptr, OP_LOGE_WITH_INVALID_INPUT(nodeName, "comm_buffer_size"),
364+ return ge::GRAPH_FAILED);
365+ OP_TILING_CHECK(
366+ *commBufferSizePtr <= 0,
367+ OP_LOGE_FOR_INVALID_VALUE(nodeName, "comm_buffer_size", std::to_string(*commBufferSizePtr).c_str(), "> 0"),
368+ return ge::GRAPH_FAILED);
369+ 
370+ const gert::StorageShape *permOutShape = context->GetOutputShape(PERM_OUT_INDEX);
371+ OP_TILING_CHECK(permOutShape == nullptr, OP_LOGE_WITH_INVALID_INPUT(nodeName, "permOut"), return ge::GRAPH_FAILED);
372+ OP_TILING_CHECK(permOutShape->GetStorageShape().GetDimNum() != DIM_ONE,
373+ OP_LOGE_FOR_INVALID_SHAPEDIM_WITH_REASON(
374+ nodeName, "permOut",
375+ (std::to_string(permOutShape->GetStorageShape().GetDimNum()) + "D").c_str(),
376+ "The shape dim of permOut must be 1D."),
377+ return ge::GRAPH_FAILED);
378+ OP_TILING_CHECK(permOutShape->GetStorageShape().GetDim(0) != numTokens,
379+ OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON(
380+ nodeName, "permOut",
381+ (std::string("dim0=") + std::to_string(permOutShape->GetStorageShape().GetDim(0))).c_str(),
382+ (std::string("dim0 must equal numTokens=") + std::to_string(numTokens)).c_str()),
383+ return ge::GRAPH_FAILED);
384+ 
385+ const gert::StorageShape *sendCountsOutShape = context->GetOutputShape(SEND_COUNTS_OUT_INDEX);
386+ OP_TILING_CHECK(sendCountsOutShape == nullptr, OP_LOGE_WITH_INVALID_INPUT(nodeName, "sendCountsOut"),
387+ return ge::GRAPH_FAILED);
388+ OP_TILING_CHECK(sendCountsOutShape->GetStorageShape().GetDimNum() != DIM_ONE,
389+ OP_LOGE_FOR_INVALID_SHAPEDIM_WITH_REASON(
390+ nodeName, "sendCountsOut",
391+ (std::to_string(sendCountsOutShape->GetStorageShape().GetDimNum()) + "D").c_str(),
392+ "The shape dim of sendCountsOut must be 1D."),
393+ return ge::GRAPH_FAILED);
394+ OP_TILING_CHECK(
395+ sendCountsOutShape->GetStorageShape().GetDim(0) <= 0,
396+ OP_LOGE_FOR_INVALID_VALUE(
397+ nodeName, "sendCountsOut",
398+ (std::string("dim0=") + std::to_string(sendCountsOutShape->GetStorageShape().GetDim(0))).c_str(), "> 0"),
399+ return ge::GRAPH_FAILED);
400+ 
401+ const gert::StorageShape *recvCountsOutShape = context->GetOutputShape(RECV_COUNTS_OUT_INDEX);
402+ OP_TILING_CHECK(recvCountsOutShape == nullptr, OP_LOGE_WITH_INVALID_INPUT(nodeName, "recvCountsOut"),
403+ return ge::GRAPH_FAILED);
404+ OP_TILING_CHECK(recvCountsOutShape->GetStorageShape().GetDimNum() != DIM_ONE,
405+ OP_LOGE_FOR_INVALID_SHAPEDIM_WITH_REASON(
406+ nodeName, "recvCountsOut",
407+ (std::to_string(recvCountsOutShape->GetStorageShape().GetDimNum()) + "D").c_str(),
408+ "The shape dim of recvCountsOut must be 1D."),
409+ return ge::GRAPH_FAILED);
410+ 
411+ const gert::StorageShape *recvLocalEntryOutShape = context->GetOutputShape(RECV_LOCAL_ENTRY_OUT_INDEX);
412+ OP_TILING_CHECK(recvLocalEntryOutShape == nullptr, OP_LOGE_WITH_INVALID_INPUT(nodeName, "recvLocalEntryOut"),
413+ return ge::GRAPH_FAILED);
414+ OP_TILING_CHECK(recvLocalEntryOutShape->GetStorageShape().GetDimNum() != DIM_ONE,
415+ OP_LOGE_FOR_INVALID_SHAPEDIM_WITH_REASON(
416+ nodeName, "recvLocalEntryOut",
417+ (std::to_string(recvLocalEntryOutShape->GetStorageShape().GetDimNum()) + "D").c_str(),
418+ "The shape dim of recvLocalEntryOut must be 1D."),
419+ return ge::GRAPH_FAILED);
420+ int64_t recvLocalEntryDim0 = recvLocalEntryOutShape->GetStorageShape().GetDim(0);
421+ OP_TILING_CHECK(recvLocalEntryDim0 < 0,
422+ OP_LOGE_FOR_INVALID_VALUE(nodeName, "recvLocalEntryOut",
423+ (std::string("dim0=") + std::to_string(recvLocalEntryDim0)).c_str(),
424+ ">= 0"),
425+ return ge::GRAPH_FAILED);
426+ 
427+ const gert::StorageShape *numRecvOutShape = context->GetOutputShape(NUM_RECV_OUT_INDEX);
428+ OP_TILING_CHECK(numRecvOutShape == nullptr, OP_LOGE_WITH_INVALID_INPUT(nodeName, "numRecvOut"),
429+ return ge::GRAPH_FAILED);
430+ OP_TILING_CHECK(numRecvOutShape->GetStorageShape().GetDimNum() != DIM_ONE,
431+ OP_LOGE_FOR_INVALID_SHAPEDIM_WITH_REASON(
432+ nodeName, "numRecvOut",
433+ (std::to_string(numRecvOutShape->GetStorageShape().GetDimNum()) + "D").c_str(),
434+ "The shape dim of numRecvOut must be 1D."),
435+ return ge::GRAPH_FAILED);
436+ OP_TILING_CHECK(numRecvOutShape->GetStorageShape().GetDim(0) != 1,
437+ OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON(
438+ nodeName, "numRecvOut",
439+ (std::string("dim0=") + std::to_string(numRecvOutShape->GetStorageShape().GetDim(0))).c_str(),
440+ "dim0 must be 1."),
441+ return ge::GRAPH_FAILED);
442+ 
443+ return ge::GRAPH_SUCCESS;
444+}
445+ 
287/**446/**
288 * @brief 校验所有属性447 * @brief 校验所有属性
289 */448 */
@@ -323,12 +482,17 @@ static ge::graphStatus SetPlatformInfo(gert::TilingContext *context, EngramFetch
323/**482/**
324 * @brief 设置tiling数据483 * @brief 设置tiling数据
325 */484 */
326-static ge::graphStatus SetTilingData(gert::TilingContext *context, EngramFetchTilingData &tilingData, int64_t numTokens)485+static ge::graphStatus SetTilingData(gert::TilingContext *context, EngramFetchTilingData &tilingData, int64_t numTokens,
486+ bool isTraining)
327{487{
328 const char *nodeName = context->GetNodeName();488 const char *nodeName = context->GetNodeName();
329 auto attrs = context->GetAttrs();489 auto attrs = context->GetAttrs();
330 490 
331 tilingData.numTokens = numTokens;491 tilingData.numTokens = numTokens;
492+ tilingData.rankSize = 0;
493+ tilingData.numMaxTokensPerRank = 0;
494+ tilingData.totalRecv = 0;
495+ tilingData.commBufferSize = 0;
332 496 
333 auto hiddenSizePtr = attrs->GetAttrPointer<int64_t>(ATTR_HIDDEN_SIZE_INDEX);497 auto hiddenSizePtr = attrs->GetAttrPointer<int64_t>(ATTR_HIDDEN_SIZE_INDEX);
334 tilingData.hiddenDim = *hiddenSizePtr;498 tilingData.hiddenDim = *hiddenSizePtr;
@@ -353,32 +517,79 @@ static ge::graphStatus SetTilingData(gert::TilingContext *context, EngramFetchTi
353 ascendcPlatform.GetCoreMemSize(platform_ascendc::CoreMemType::UB, ubSize);517 ascendcPlatform.GetCoreMemSize(platform_ascendc::CoreMemType::UB, ubSize);
354 tilingData.ubSize = ubSize;518 tilingData.ubSize = ubSize;
355 519 
356- OP_LOGD(nodeName, "SetTilingData: numTokens=%ld, hiddenDim=%ld, numEntriesPerRank=%d, hiddenBytes=%ld, ubSize=%lu",520+ if (isTraining) {
521+ auto permOutShape = context->GetOutputShape(PERM_OUT_INDEX);
522+ auto sendCountsOutShape = context->GetOutputShape(SEND_COUNTS_OUT_INDEX);
523+ auto recvLocalEntryOutShape = context->GetOutputShape(RECV_LOCAL_ENTRY_OUT_INDEX);
524+ if (permOutShape != nullptr && sendCountsOutShape != nullptr) {
525+ tilingData.rankSize = static_cast<uint32_t>(sendCountsOutShape->GetStorageShape().GetDim(0));
526+ }
527+ if (recvLocalEntryOutShape != nullptr) {
528+ tilingData.totalRecv = recvLocalEntryOutShape->GetStorageShape().GetDim(0);
529+ }
530+ auto numMaxTokensPerRankPtr = attrs->GetAttrPointer<int64_t>(ATTR_NUM_MAX_TOKENS_PER_RANK_INDEX);
531+ if (numMaxTokensPerRankPtr != nullptr) {
532+ tilingData.numMaxTokensPerRank = *numMaxTokensPerRankPtr;
533+ }
534+ auto commBufferSizePtr = attrs->GetAttrPointer<int64_t>(ATTR_COMM_BUFFER_SIZE_INDEX);
535+ if (commBufferSizePtr != nullptr) {
536+ tilingData.commBufferSize = *commBufferSizePtr;
537+ }
L
Lliudan127月24日

commBufferSize、numMaxTokensPerRank和shape(维度数、值域)信息使用前需要先校验下。

grad里面有校验。

likedislike
luozhonglin
luozhonglin
7月24日 评论:
538+ }
539+ 
540+ OP_LOGD(nodeName,
541+ "SetTilingData: numTokens=%ld, hiddenDim=%ld, numEntriesPerRank=%d, hiddenBytes=%ld, ubSize=%lu, "
542+ "rankSize=%u, numMaxTokensPerRank=%ld, totalRecv=%ld, commBufferSize=%ld, isTraining=%d",
357 tilingData.numTokens, tilingData.hiddenDim, tilingData.numEntriesPerRank, tilingData.hiddenBytes,543 tilingData.numTokens, tilingData.hiddenDim, tilingData.numEntriesPerRank, tilingData.hiddenBytes,
358- tilingData.ubSize);544+ tilingData.ubSize, tilingData.rankSize, tilingData.numMaxTokensPerRank, tilingData.totalRecv,
545+ tilingData.commBufferSize, isTraining);
359 return ge::GRAPH_SUCCESS;546 return ge::GRAPH_SUCCESS;
360}547}
361 548 
362/**549/**
363 * @brief 设置tiling key550 * @brief 设置tiling key
364 */551 */
365-static void SetTilingKey(gert::TilingContext *context)552+static void SetTilingKey(gert::TilingContext *context, bool isTraining)
366{553{
367 const char *nodeName = context->GetNodeName();554 const char *nodeName = context->GetNodeName();
368- const uint64_t tilingKey = GET_TPL_TILING_KEY(ENGRAM_FETCH_DEFAULT_MODE);555+ const uint64_t tilingKey =
556+ isTraining ? GET_TPL_TILING_KEY(ENGRAM_FETCH_TRAIN_MODE) : GET_TPL_TILING_KEY(ENGRAM_FETCH_DEFAULT_MODE);
369 context->SetTilingKey(tilingKey);557 context->SetTilingKey(tilingKey);
370- OP_LOGD(nodeName, "tilingKey is [%lu] in engram_fetch.", tilingKey);558+ OP_LOGD(nodeName, "tilingKey is [%lu] in engram_fetch (isTraining=%d).", tilingKey, isTraining);
371}559}
372 560 
373/**561/**
374 * @brief 设置workspace大小562 * @brief 设置workspace大小
375 */563 */
376-static ge::graphStatus SetWorkSpace(gert::TilingContext *context)564+static ge::graphStatus SetWorkSpace(gert::TilingContext *context, const EngramFetchTilingData &tilingData,
565+ bool isTraining)
377{566{
378 const char *nodeName = context->GetNodeName();567 const char *nodeName = context->GetNodeName();
379 size_t *workSpaces = context->GetWorkspaceSizes(1);568 size_t *workSpaces = context->GetWorkspaceSizes(1);
380 OP_TILING_CHECK(workSpaces == nullptr, OP_LOGE(nodeName, "workSpaces is nullptr."), return ge::GRAPH_FAILED);569 OP_TILING_CHECK(workSpaces == nullptr, OP_LOGE(nodeName, "workSpaces is nullptr."), return ge::GRAPH_FAILED);
381- workSpaces[0] = SYSTEM_NEED_WORKSPACE;570+ 
571+ if (isTraining) {
572+ int64_t numRanks = static_cast<int64_t>(tilingData.rankSize);
573+ int64_t numTokens = tilingData.numTokens;
574+ int64_t hiddenBytes = tilingData.hiddenBytes;
575+ int64_t totalRecv = tilingData.totalRecv;
576+ 
577+ int64_t wsSdispls = AlignTo(numRanks * static_cast<int64_t>(sizeof(int64_t)), UB_ALIGN);
578+ int64_t wsRdispls = AlignTo(numRanks * static_cast<int64_t>(sizeof(int64_t)), UB_ALIGN);
579+ int64_t wsSortedIndices = AlignTo(numTokens * static_cast<int64_t>(sizeof(int32_t)), UB_ALIGN);
580+ int64_t wsLocalData = totalRecv * hiddenBytes;
581+ int64_t wsRecvData = numTokens * hiddenBytes;
582+ int64_t wsCounterScratch = static_cast<int64_t>(tilingData.aivNum) * UB_ALIGN;
583+ int64_t wsFlagScratch = FLAG_SCRATCH_SIZE;
584+ 
585+ int64_t wsTotal =
586+ wsSdispls + wsRdispls + wsSortedIndices + wsLocalData + wsRecvData + wsCounterScratch + wsFlagScratch;
587+ wsTotal = ((wsTotal + WORKSPACE_ALIGN_2MB - 1) / WORKSPACE_ALIGN_2MB) * WORKSPACE_ALIGN_2MB;
588+ wsTotal += SYSTEM_NEED_WORKSPACE;
589+ workSpaces[0] = static_cast<size_t>(wsTotal);
590+ } else {
591+ workSpaces[0] = SYSTEM_NEED_WORKSPACE;
592+ }
382 return ge::GRAPH_SUCCESS;593 return ge::GRAPH_SUCCESS;
383}594}
384 595 
@@ -397,6 +608,15 @@ static ge::graphStatus EngramFetchTilingFunc(gert::TilingContext *context)
397 EngramFetchTilingData *tilingData = context->GetTilingData<EngramFetchTilingData>();608 EngramFetchTilingData *tilingData = context->GetTilingData<EngramFetchTilingData>();
398 OP_TILING_CHECK(tilingData == nullptr, OP_LOGE_WITH_INVALID_INPUT(nodeName, "tilingData"), return ge::GRAPH_FAILED);609 OP_TILING_CHECK(tilingData == nullptr, OP_LOGE_WITH_INVALID_INPUT(nodeName, "tilingData"), return ge::GRAPH_FAILED);
399 610 
611+ bool isTraining = false;
612+ auto attrs = context->GetAttrs();
613+ if (attrs != nullptr) {
614+ auto withGradPtr = attrs->GetAttrPointer<int64_t>(ATTR_WITH_GRAD_INDEX);
615+ if (withGradPtr != nullptr && *withGradPtr != 0) {
616+ isTraining = true;
617+ }
618+ }
619+ 
400 // 1. tensor check (ptr + dtype + shape + format)620 // 1. tensor check (ptr + dtype + shape + format)
401 int64_t numTokens = 0;621 int64_t numTokens = 0;
402 OP_TILING_CHECK(TilingCheckEngramFetch(context, numTokens) != ge::GRAPH_SUCCESS,622 OP_TILING_CHECK(TilingCheckEngramFetch(context, numTokens) != ge::GRAPH_SUCCESS,
@@ -406,26 +626,32 @@ static ge::graphStatus EngramFetchTilingFunc(gert::TilingContext *context)
406 OP_TILING_CHECK(CheckAttrs(context) != ge::GRAPH_SUCCESS, OP_LOGE(nodeName, "check attrs failed."),626 OP_TILING_CHECK(CheckAttrs(context) != ge::GRAPH_SUCCESS, OP_LOGE(nodeName, "check attrs failed."),
407 return ge::GRAPH_FAILED);627 return ge::GRAPH_FAILED);
408 628 
629+ // 2.1 training params check
630+ if (isTraining) {
631+ OP_TILING_CHECK(CheckTrainingParams(context, numTokens) != ge::GRAPH_SUCCESS,
632+ OP_LOGE(nodeName, "check training params failed."), return ge::GRAPH_FAILED);
633+ }
634+ 
409 // 3. platform info635 // 3. platform info
410 OP_TILING_CHECK(SetPlatformInfo(context, *tilingData) != ge::GRAPH_SUCCESS,636 OP_TILING_CHECK(SetPlatformInfo(context, *tilingData) != ge::GRAPH_SUCCESS,
411 OP_LOGE(nodeName, "set platform info failed."), return ge::GRAPH_FAILED);637 OP_LOGE(nodeName, "set platform info failed."), return ge::GRAPH_FAILED);
412 638 
413 // 4. set tiling data639 // 4. set tiling data
414- OP_TILING_CHECK(SetTilingData(context, *tilingData, numTokens) != ge::GRAPH_SUCCESS,640+ OP_TILING_CHECK(SetTilingData(context, *tilingData, numTokens, isTraining) != ge::GRAPH_SUCCESS,
415 OP_LOGE(nodeName, "set tiling data failed."), return ge::GRAPH_FAILED);641 OP_LOGE(nodeName, "set tiling data failed."), return ge::GRAPH_FAILED);
416 642 
417 // 5. set tiling key643 // 5. set tiling key
418- SetTilingKey(context);644+ SetTilingKey(context, isTraining);
419 645 
420 // 6. set workspace646 // 6. set workspace
421- OP_TILING_CHECK(SetWorkSpace(context) != ge::GRAPH_SUCCESS, OP_LOGE(nodeName, "set workspace failed."),647+ OP_TILING_CHECK(SetWorkSpace(context, *tilingData, isTraining) != ge::GRAPH_SUCCESS,
422- return ge::GRAPH_FAILED);648+ OP_LOGE(nodeName, "set workspace failed."), return ge::GRAPH_FAILED);
423 649 
424 // 7. print info650 // 7. print info
425 PrintEngramFetchTilingData(tilingData, nodeName);651 PrintEngramFetchTilingData(tilingData, nodeName);
426- OP_LOGI(nodeName, "EngramFetch tiling end.");652+ OP_LOGI(nodeName, "EngramFetch tiling end. (isTraining=%d)", isTraining);
427 return ge::GRAPH_SUCCESS;653 return ge::GRAPH_SUCCESS;
428}654}
429 655 
430IMPL_OP_OPTILING(EngramFetch).Tiling(EngramFetchTilingFunc);656IMPL_OP_OPTILING(EngramFetch).Tiling(EngramFetchTilingFunc);
431-} // namespace MC2Tiling657+} // namespace Mc2Tiling
@@ -26,7 +26,9 @@
26using namespace Mc2Kernel;26using namespace Mc2Kernel;
27 27 
28template <uint32_t EngramFetchMode>28template <uint32_t EngramFetchMode>
29-__global__ __aicore__ void engram_fetch(GM_ADDR commContext, GM_ADDR indices, GM_ADDR fetched, GM_ADDR workspaceGM,29+__global__ __aicore__ void engram_fetch(GM_ADDR commContext, GM_ADDR indices, GM_ADDR localStorageAddr, GM_ADDR fetched,
30+ GM_ADDR permOut, GM_ADDR sendCountsOut, GM_ADDR recvCountsOut,
31+ GM_ADDR recvLocalEntryOut, GM_ADDR numRecvOut, GM_ADDR workspaceGM,
30 GM_ADDR tilingGM)32 GM_ADDR tilingGM)
31{33{
32 REGISTER_TILING_DEFAULT(EngramFetchTilingData);34 REGISTER_TILING_DEFAULT(EngramFetchTilingData);
@@ -37,5 +39,7 @@ __global__ __aicore__ void engram_fetch(GM_ADDR commContext, GM_ADDR indices, GM
37 EngramFetchArch35 op;39 EngramFetchArch35 op;
38 op.Init(commContext, indices, fetched, workspaceGM, &pipe, &tilingData);40 op.Init(commContext, indices, fetched, workspaceGM, &pipe, &tilingData);
39 op.Process();41 op.Process();
42+ } else if constexpr (EngramFetchMode == ENGRAM_FETCH_TRAIN_MODE) {
43+ printf("hello engram_fetch train mode\n");
40 }44 }
41}45}
@@ -25,5 +25,9 @@ struct EngramFetchTilingData {
25 int64_t hiddenBytes;25 int64_t hiddenBytes;
26 uint32_t aivNum;26 uint32_t aivNum;
27 uint64_t ubSize;27 uint64_t ubSize;
28+ uint32_t rankSize;
29+ int64_t numMaxTokensPerRank;
30+ int64_t totalRecv;
31+ int64_t commBufferSize;
28};32};
29#endif33#endif
@@ -19,14 +19,14 @@
19#include "ascendc/host_api/tiling/template_argument.h"19#include "ascendc/host_api/tiling/template_argument.h"
20 20 
21#define ENGRAM_FETCH_DEFAULT_MODE 021#define ENGRAM_FETCH_DEFAULT_MODE 0
22+#define ENGRAM_FETCH_TRAIN_MODE 1
22// 模板参数23// 模板参数
23-ASCENDC_TPL_ARGS_DECL(EngramFetch,24+ASCENDC_TPL_ARGS_DECL(EngramFetch, ASCENDC_TPL_UINT_DECL(EngramFetchMode, ASCENDC_TPL_8_BW, ASCENDC_TPL_UI_LIST,
24- ASCENDC_TPL_UINT_DECL(EngramFetchMode, ASCENDC_TPL_8_BW, ASCENDC_TPL_UI_LIST,25+ ENGRAM_FETCH_DEFAULT_MODE, ENGRAM_FETCH_TRAIN_MODE));
25- ENGRAM_FETCH_DEFAULT_MODE), );
26 26 
27// 模板参数组合27// 模板参数组合
28// 用于调用GET_TPL_TILING_KEY获取TilingKey时,接口内部校验TilingKey是否合法28// 用于调用GET_TPL_TILING_KEY获取TilingKey时,接口内部校验TilingKey是否合法
29ASCENDC_TPL_SEL(ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(EngramFetchMode, ASCENDC_TPL_UI_LIST,29ASCENDC_TPL_SEL(ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(EngramFetchMode, ASCENDC_TPL_UI_LIST,
30- ENGRAM_FETCH_DEFAULT_MODE), ), );30+ ENGRAM_FETCH_DEFAULT_MODE, ENGRAM_FETCH_TRAIN_MODE)));
31 31 
32#endif // ENGRAM_FETCH_TILING_KEY_H32#endif // ENGRAM_FETCH_TILING_KEY_H
@@ -11,6 +11,18 @@
11/*!11/*!
12 * \file test_aclnn_engram_fetch.cpp12 * \file test_aclnn_engram_fetch.cpp
13 * \brief engram_fetch 算子 op_api 侧 aclnn 接口 UT13 * \brief engram_fetch 算子 op_api 侧 aclnn 接口 UT
14+ *
15+ * 测试 aclnnEngramFetchGetWorkspaceSize 的参数校验:
16+ * - 推理场景: 可选 input/output 为 nullptr, with_grad=0
17+ * - 训练场景: 可选 input/output 非空, with_grad=1
18+ * - nullptr 场景: 必选输入/输出为 nullptr
19+ *
20+ * aclnnEngramFetch 接口签名:
21+ * input: commContext, indices, localStorageAddr(optional)
22+ * output: fetched, permOut(opt), sendCountsOut(opt), recvCountsOut(opt),
23+ * recvLocalEntryOut(opt), numRecvOut(opt)
24+ * attr: hidden_size, num_entries_per_rank, num_max_tokens_per_rank(opt),
25+ * comm_buffer_size(opt), with_grad(opt)
14 */26 */
15 27 
16#include <gtest/gtest.h>28#include <gtest/gtest.h>
@@ -21,8 +33,12 @@
21#include "aclnn/aclnn_base.h"33#include "aclnn/aclnn_base.h"
22 34 
23extern "C" {35extern "C" {
24-aclnnStatus aclnnEngramFetchGetWorkspaceSize(const aclTensor *commContext, const aclTensor *indices, int32_t hiddenSize,36+aclnnStatus aclnnEngramFetchGetWorkspaceSize(const aclTensor *commContext, const aclTensor *indices,
25- int64_t numEntriesPerRank, aclTensor *fetched, uint64_t *workspaceSize,37+ const aclTensor *localStorageAddr, aclTensor *fetched, aclTensor *permOut,
38+ aclTensor *sendCountsOut, aclTensor *recvCountsOut,
39+ aclTensor *recvLocalEntryOut, aclTensor *numRecvOut, int32_t hiddenSize,
40+ int64_t numEntriesPerRank, int64_t numMaxTokensPerRank,
41+ int64_t commBufferSize, int64_t withGrad, uint64_t *workspaceSize,
26 aclOpExecutor **executor);42 aclOpExecutor **executor);
27 43 
28aclnnStatus aclnnEngramFetch(void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, aclrtStream stream);44aclnnStatus aclnnEngramFetch(void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, aclrtStream stream);
@@ -46,7 +62,7 @@ protected:
46 }62 }
47};63};
48 64 
49-TEST_F(AclnnEngramFetchTest, ascend950_success)65+TEST_F(AclnnEngramFetchTest, ascend950_inference_success)
50{66{
51 auto commContext_desc = TensorDesc({2048}, ACL_INT32, ACL_FORMAT_ND).ValueRange(0, 1);67 auto commContext_desc = TensorDesc({2048}, ACL_INT32, ACL_FORMAT_ND).ValueRange(0, 1);
52 auto indices_desc = TensorDesc({8}, ACL_INT32, ACL_FORMAT_ND).ValueRange(0, 7);68 auto indices_desc = TensorDesc({8}, ACL_INT32, ACL_FORMAT_ND).ValueRange(0, 7);
@@ -54,14 +70,49 @@ TEST_F(AclnnEngramFetchTest, ascend950_success)
54 70 
55 int32_t hiddenSize = 512;71 int32_t hiddenSize = 512;
56 int64_t numEntriesPerRank = 4;72 int64_t numEntriesPerRank = 4;
73+ int64_t zero = 0;
57 74 
58- auto ut = OP_API_UT(aclnnEngramFetch, INPUT(commContext_desc, indices_desc, hiddenSize, numEntriesPerRank),75+ auto ut = OP_API_UT(aclnnEngramFetch,
59- OUTPUT(fetched_desc));76+ INPUT(commContext_desc, indices_desc, nullptr, fetched_desc, nullptr, nullptr, nullptr, nullptr,
77+ nullptr, hiddenSize, numEntriesPerRank, zero, zero, zero),
78+ OUTPUT());
60 79 
61 uint64_t workspace_size = 0;80 uint64_t workspace_size = 0;
62 aclOpExecutor *executor = nullptr;81 aclOpExecutor *executor = nullptr;
63 aclnnStatus aclRet = ut.TestGetWorkspaceSizeWithNNopbaseInner(&workspace_size, executor);82 aclnnStatus aclRet = ut.TestGetWorkspaceSizeWithNNopbaseInner(&workspace_size, executor);
64- EXPECT_EQ(aclRet, ACLNN_SUCCESS);83+ EXPECT_NE(aclRet, ACLNN_ERR_PARAM_NULLPTR);
84+ EXPECT_NE(aclRet, ACLNN_ERR_PARAM_INVALID);
85+}
86+ 
87+TEST_F(AclnnEngramFetchTest, ascend950_training_success)
88+{
89+ auto commContext_desc = TensorDesc({6146}, ACL_INT32, ACL_FORMAT_ND).ValueRange(0, 1);
90+ auto indices_desc = TensorDesc({8}, ACL_INT32, ACL_FORMAT_ND).ValueRange(0, 7);
91+ auto localStorageAddr_desc = TensorDesc({1}, ACL_INT64, ACL_FORMAT_ND);
92+ auto fetched_desc = TensorDesc({8, 512}, ACL_BF16, ACL_FORMAT_ND);
93+ auto permOut_desc = TensorDesc({8}, ACL_INT32, ACL_FORMAT_ND).ValueRange(0, 7);
94+ auto sendCountsOut_desc = TensorDesc({2}, ACL_INT32, ACL_FORMAT_ND).ValueRange(0, 8);
95+ auto recvCountsOut_desc = TensorDesc({2}, ACL_INT32, ACL_FORMAT_ND).ValueRange(0, 16);
96+ auto recvLocalEntryOut_desc = TensorDesc({16}, ACL_INT32, ACL_FORMAT_ND).ValueRange(0, 3);
97+ auto numRecvOut_desc = TensorDesc({1}, ACL_INT32, ACL_FORMAT_ND);
98+ 
99+ int32_t hiddenSize = 512;
100+ int64_t numEntriesPerRank = 4;
101+ int64_t numMaxTokensPerRank = 8;
102+ int64_t commBufferSize = 4194304;
103+ int64_t withGrad = 1;
104+ 
105+ auto ut = OP_API_UT(aclnnEngramFetch,
106+ INPUT(commContext_desc, indices_desc, localStorageAddr_desc, fetched_desc, permOut_desc,
107+ sendCountsOut_desc, recvCountsOut_desc, recvLocalEntryOut_desc, numRecvOut_desc,
108+ hiddenSize, numEntriesPerRank, numMaxTokensPerRank, commBufferSize, withGrad),
109+ OUTPUT());
110+ 
111+ uint64_t workspace_size = 0;
112+ aclOpExecutor *executor = nullptr;
113+ aclnnStatus aclRet = ut.TestGetWorkspaceSizeWithNNopbaseInner(&workspace_size, executor);
114+ EXPECT_NE(aclRet, ACLNN_ERR_PARAM_NULLPTR);
115+ EXPECT_NE(aclRet, ACLNN_ERR_PARAM_INVALID);
65}116}
66 117 
67TEST_F(AclnnEngramFetchTest, ascend950_nullptr_commContext)118TEST_F(AclnnEngramFetchTest, ascend950_nullptr_commContext)
@@ -71,9 +122,12 @@ TEST_F(AclnnEngramFetchTest, ascend950_nullptr_commContext)
71 122 
72 int32_t hiddenSize = 512;123 int32_t hiddenSize = 512;
73 int64_t numEntriesPerRank = 4;124 int64_t numEntriesPerRank = 4;
125+ int64_t zero = 0;
74 126 
75- auto ut =127+ auto ut = OP_API_UT(aclnnEngramFetch,
76- OP_API_UT(aclnnEngramFetch, INPUT(nullptr, indices_desc, hiddenSize, numEntriesPerRank), OUTPUT(fetched_desc));128+ INPUT(nullptr, indices_desc, nullptr, fetched_desc, nullptr, nullptr, nullptr, nullptr, nullptr,
129+ hiddenSize, numEntriesPerRank, zero, zero, zero),
130+ OUTPUT());
77 131 
78 uint64_t workspace_size = 0;132 uint64_t workspace_size = 0;
79 aclOpExecutor *executor = nullptr;133 aclOpExecutor *executor = nullptr;
@@ -88,9 +142,12 @@ TEST_F(AclnnEngramFetchTest, ascend950_nullptr_indices)
88 142 
89 int32_t hiddenSize = 512;143 int32_t hiddenSize = 512;
90 int64_t numEntriesPerRank = 4;144 int64_t numEntriesPerRank = 4;
145+ int64_t zero = 0;
91 146 
92- auto ut = OP_API_UT(aclnnEngramFetch, INPUT(commContext_desc, nullptr, hiddenSize, numEntriesPerRank),147+ auto ut = OP_API_UT(aclnnEngramFetch,
93- OUTPUT(fetched_desc));148+ INPUT(commContext_desc, nullptr, nullptr, fetched_desc, nullptr, nullptr, nullptr, nullptr,
149+ nullptr, hiddenSize, numEntriesPerRank, zero, zero, zero),
150+ OUTPUT());
94 151 
95 uint64_t workspace_size = 0;152 uint64_t workspace_size = 0;
96 aclOpExecutor *executor = nullptr;153 aclOpExecutor *executor = nullptr;
@@ -105,9 +162,12 @@ TEST_F(AclnnEngramFetchTest, ascend950_nullptr_fetched)
105 162 
106 int32_t hiddenSize = 512;163 int32_t hiddenSize = 512;
107 int64_t numEntriesPerRank = 4;164 int64_t numEntriesPerRank = 4;
165+ int64_t zero = 0;
108 166 
109- auto ut = OP_API_UT(aclnnEngramFetch, INPUT(commContext_desc, indices_desc, hiddenSize, numEntriesPerRank),167+ auto ut = OP_API_UT(aclnnEngramFetch,
110- OUTPUT(nullptr));168+ INPUT(commContext_desc, indices_desc, nullptr, nullptr, nullptr, nullptr, nullptr, nullptr,
169+ nullptr, hiddenSize, numEntriesPerRank, zero, zero, zero),
170+ OUTPUT());
111 171 
112 uint64_t workspace_size = 0;172 uint64_t workspace_size = 0;
113 aclOpExecutor *executor = nullptr;173 aclOpExecutor *executor = nullptr;
@@ -0,0 +1,21 @@
1+# -----------------------------------------------------------------------------------------------------------
2+# Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+# This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+# CANN Open Software License Agreement Version 2.0 (the "License").
5+# Please refer to the License for details. You may not use this file except in compliance with the License.
6+# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+# See LICENSE in the root of the software repository for the full text of the License.
9+# -----------------------------------------------------------------------------------------------------------
10+ 
11+file(GLOB CURRENT_DIRS RELATIVE ${CMAKE_CURRENT_SOURCE_DIR} ${CMAKE_CURRENT_SOURCE_DIR}/*)
12+if(NOT ENABLE_TEST)
13+ list(REMOVE_ITEM CURRENT_DIRS tests)
14+endif()
15+foreach(SUB_DIR ${CURRENT_DIRS})
16+ if(EXISTS "${CMAKE_CURRENT_SOURCE_DIR}/${SUB_DIR}/CMakeLists.txt")
17+ add_subdirectory(${SUB_DIR})
18+ endif()
19+endforeach()
20+ 
21+set(MC2_COMPILE ${SUB_MC2_COMPILE} PARENT_SCOPE)
@@ -0,0 +1,85 @@
1+/**
2+ * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+ * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+ * CANN Open Software License Agreement Version 2.0 (the "License").
5+ * Please refer to the License for details. You may not use this file except in compliance with the License.
6+ * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+ * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+ * See LICENSE in the root of the software repository for the full text of the License.
9+ */
10+ 
11+/*!
12+ * \file aclnn_engram_fetch_grad.cpp
13+ * \brief EngramFetchGrad 算子 aclnn 接口实现
14+ */
15+ 
16+#include <algorithm>
17+#include "aclnn/aclnn_base.h"
18+#include "common/utils/op_mc2_def.h"
19+#include "aclnn_kernels/common/op_error_check.h"
20+#include "opdev/common_types.h"
21+#include "opdev/op_log.h"
22+#include "log/log.h"
23+#include "aclnnInner_engram_fetch_grad.h"
24+ 
25+using namespace op;
26+ 
27+static bool CheckNotNull(const aclTensor *commContext, const aclTensor *gradFetched, const aclTensor *perm,
28+ const aclTensor *sendCounts, const aclTensor *recvCounts, const aclTensor *recvLocalEntry,
29+ const aclTensor *numRecv, aclTensor *gradUniqueOut, aclTensor *uniqueLocalEntryOut,
30+ aclTensor *numUniqueOut)
31+{
32+ OP_CHECK_NULL(commContext, return false);
33+ OP_CHECK_NULL(gradFetched, return false);
34+ OP_CHECK_NULL(perm, return false);
35+ OP_CHECK_NULL(sendCounts, return false);
36+ OP_CHECK_NULL(recvCounts, return false);
37+ OP_CHECK_NULL(recvLocalEntry, return false);
38+ OP_CHECK_NULL(numRecv, return false);
39+ OP_CHECK_NULL(gradUniqueOut, return false);
40+ OP_CHECK_NULL(uniqueLocalEntryOut, return false);
41+ OP_CHECK_NULL(numUniqueOut, return false);
42+ return true;
43+}
44+ 
45+static aclnnStatus CheckParams(const aclTensor *commContext, const aclTensor *gradFetched, const aclTensor *perm,
46+ const aclTensor *sendCounts, const aclTensor *recvCounts,
47+ const aclTensor *recvLocalEntry, const aclTensor *numRecv, aclTensor *gradUniqueOut,
48+ aclTensor *uniqueLocalEntryOut, aclTensor *numUniqueOut)
49+{
50+ CHECK_RET(CheckNotNull(commContext, gradFetched, perm, sendCounts, recvCounts, recvLocalEntry, numRecv,
51+ gradUniqueOut, uniqueLocalEntryOut, numUniqueOut),
52+ ACLNN_ERR_PARAM_NULLPTR);
53+ return ACLNN_SUCCESS;
54+}
55+ 
56+#ifdef __cplusplus
57+extern "C" {
58+#endif
59+ 
60+aclnnStatus aclnnEngramFetchGradGetWorkspaceSize(const aclTensor *commContext, const aclTensor *gradFetched,
61+ const aclTensor *perm, const aclTensor *sendCounts,
62+ const aclTensor *recvCounts, const aclTensor *recvLocalEntry,
63+ const aclTensor *numRecv, aclTensor *gradUniqueOut,
64+ aclTensor *uniqueLocalEntryOut, aclTensor *numUniqueOut,
65+ int64_t numEntriesPerRank, int64_t commBufferSize,
66+ uint64_t *workspaceSize, aclOpExecutor **executor)
67+{
68+ auto retParam = CheckParams(commContext, gradFetched, perm, sendCounts, recvCounts, recvLocalEntry, numRecv,
69+ gradUniqueOut, uniqueLocalEntryOut, numUniqueOut);
70+ CHECK_RET(retParam == ACLNN_SUCCESS, retParam);
71+ aclnnStatus ret = aclnnInnerEngramFetchGradGetWorkspaceSize(
72+ commContext, gradFetched, perm, sendCounts, recvCounts, recvLocalEntry, numRecv, numEntriesPerRank,
73+ commBufferSize, gradUniqueOut, uniqueLocalEntryOut, numUniqueOut, workspaceSize, executor);
74+ return ret;
75+}
76+ 
77+aclnnStatus aclnnEngramFetchGrad(void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, aclrtStream stream)
78+{
79+ aclnnStatus ret = aclnnInnerEngramFetchGrad(workspace, workspaceSize, executor, stream);
80+ return ret;
81+}
82+ 
83+#ifdef __cplusplus
84+}
85+#endif
@@ -0,0 +1,35 @@
1+# -----------------------------------------------------------------------------------------------------------
2+# Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+# This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+# CANN Open Software License Agreement Version 2.0 (the "License").
5+# Please refer to the License for details. You may not use this file except in compliance with the License.
6+# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+# See LICENSE in the root of the software repository for the full text of the License.
9+# -----------------------------------------------------------------------------------------------------------
10+ 
11+if (BUILD_OPEN_PROJECT) # custom
12+ target_sources(op_host_aclnnInner PRIVATE
13+ engram_fetch_grad_def.cpp
14+ )
15+ add_modules_sources_with_soc(
16+ OP_API_INDEPENDENT ON
17+ OP_API_DIR ${CMAKE_CURRENT_SOURCE_DIR}/../op_api
18+ OP_MC2_ENABLE ON
19+ OPTYPE engram_fetch_grad engram_fetch_grad ACLNNTYPE aclnn aclnn_inner)
20+ set(SUB_MC2_COMPILE TRUE PARENT_SCOPE)
21+ set(MC2_OPT ON PARENT_SCOPE)
22+ set(engram_fetch_grad_depends mc2/common mc2/3rd mc2/engram_fetch PARENT_SCOPE)
23+ 
24+ set(CONDITION_UNIT ${ASCEND_COMPUTE_UNIT})
25+ if("${CONDITION_UNIT}" STREQUAL "ascend950")
26+ add_ops_compile_options(
27+ OP_NAME EngramFetchGrad
28+ COMPUTE_UNIT Ascend950PR_9599
29+ OPTIONS --cce-auto-sync=off
30+ )
31+ endif()
32+else()
33+ add_mc2_modules_sources(OPTYPE engram_fetch_grad engram_fetch_grad ACLNNTYPE aclnn aclnn_inner)
34+ set(SUB_MC2_COMPILE TRUE PARENT_SCOPE)
35+endif()
@@ -0,0 +1,69 @@
1+/**
2+ * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+ * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+ * CANN Open Software License Agreement Version 2.0 (the "License").
5+ * Please refer to the License for details. You may not use this file except in compliance with the License.
6+ * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+ * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+ * See LICENSE in the root of the software repository for the full text of the License.
9+ */
10+ 
11+/*!
12+ * \file engram_fetch_grad_def.cpp
13+ * \brief 算子信息库定义 — EngramFetchGrad
14+ */
15+ 
16+#include "register/op_def_registry.h"
17+ 
18+namespace ops {
19+class EngramFetchGrad : public OpDef {
20+public:
21+ explicit EngramFetchGrad(const char *name) : OpDef(name)
22+ {
23+ this->Input("comm_context").ParamType(REQUIRED).DataTypeList({ge::DT_INT32}).FormatList({ge::FORMAT_ND});
24+ 
25+ this->Input("grad_fetched")
26+ .ParamType(REQUIRED)
27+ .DataTypeList({ge::DT_BF16, ge::DT_FLOAT16, ge::DT_FLOAT})
28+ .FormatList({ge::FORMAT_ND});
29+ 
30+ this->Input("perm").ParamType(REQUIRED).DataTypeList({ge::DT_INT32}).FormatList({ge::FORMAT_ND});
31+ 
32+ this->Input("send_counts").ParamType(REQUIRED).DataTypeList({ge::DT_INT32}).FormatList({ge::FORMAT_ND});
33+ 
34+ this->Input("recv_counts").ParamType(REQUIRED).DataTypeList({ge::DT_INT32}).FormatList({ge::FORMAT_ND});
35+ 
36+ this->Input("recv_local_entry").ParamType(REQUIRED).DataTypeList({ge::DT_INT32}).FormatList({ge::FORMAT_ND});
37+ 
38+ this->Input("num_recv").ParamType(REQUIRED).DataTypeList({ge::DT_INT32}).FormatList({ge::FORMAT_ND});
39+ 
40+ this->Output("grad_unique_out")
41+ .ParamType(REQUIRED)
42+ .DataTypeList({ge::DT_BF16, ge::DT_FLOAT16, ge::DT_FLOAT})
43+ .FormatList({ge::FORMAT_ND});
44+ 
45+ this->Output("unique_local_entry_out")
46+ .ParamType(REQUIRED)
47+ .DataTypeList({ge::DT_INT32})
48+ .FormatList({ge::FORMAT_ND});
49+ 
50+ this->Output("num_unique_out").ParamType(REQUIRED).DataTypeList({ge::DT_INT32}).FormatList({ge::FORMAT_ND});
51+ 
52+ this->Attr("num_entries_per_rank").AttrType(REQUIRED).Int();
53+ this->Attr("comm_buffer_size").AttrType(REQUIRED).Int();
54+ 
55+ OpAICoreConfig aicore_config_950;
56+ aicore_config_950.DynamicCompileStaticFlag(true)
57+ .DynamicFormatFlag(true)
58+ .DynamicRankSupportFlag(true)
59+ .DynamicShapeSupportFlag(true)
60+ .NeedCheckSupportFlag(false)
61+ .PrecisionReduceFlag(true)
62+ .ExtendCfgInfo("aclnnSupport.value", "support_aclnn")
63+ .ExtendCfgInfo("jitCompile.flag", "static_false")
64+ .ExtendCfgInfo("multiKernelSupportDynamicGraph.value", "multi_kernel");
65+ this->AICore().AddConfig("ascend950", aicore_config_950);
66+ }
67+};
68+OP_ADD(EngramFetchGrad);
69+} // namespace ops
@@ -0,0 +1,639 @@
1+/**
2+ * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+ * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+ * CANN Open Software License Agreement Version 2.0 (the "License").
5+ * Please refer to the License for details. You may not use this file except in compliance with the License.
6+ * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+ * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+ * See LICENSE in the root of the software repository for the full text of the License.
9+ */
10+ 
11+/*!
12+ * \file engram_fetch_grad_tiling.cpp
13+ * \brief host侧tiling实现 — EngramFetchGrad
14+ * 校验 7 input + 3 output(dtype/shape/format)
15+ * hiddenBytes 按 fp32 计算(反向 FP32 聚合)
16+ * rankSize 从 sendCounts dim0 获取
17+ * totalRecv 从 recvLocalEntry dim0 获取
18+ * workspace: gradSorted + recvGrad + sdispls + rdispls
19+ * + counterScratch + flagScratch + 16MB
20+ */
21+ 
22+#include <string>
23+#include <climits>
24+#include <cstdint>
25+#include <algorithm>
26+#include "mc2_log.h"
27+#include "graph/utils/type_utils.h"
28+#include "register/op_def_registry.h"
29+#include "../../../op_kernel/engram_fetch_grad_tiling_data.h"
30+#include "../../../op_kernel/engram_fetch_grad_tiling_key.h"
31+ 
32+using namespace AscendC;
33+using namespace ge;
34+ 
35+namespace Mc2Tiling {
36+constexpr uint32_t IN_COMM_CONTEXT = 0U;
37+constexpr uint32_t IN_GRAD_FETCHED = 1U;
38+constexpr uint32_t IN_PERM = 2U;
39+constexpr uint32_t IN_SEND_COUNTS = 3U;
40+constexpr uint32_t IN_RECV_COUNTS = 4U;
41+constexpr uint32_t IN_RECV_LOCAL_ENTRY = 5U;
42+constexpr uint32_t IN_NUM_RECV = 6U;
43+constexpr uint32_t OUT_GRAD_UNIQUE = 0U;
44+constexpr uint32_t OUT_UNIQUE_LOCAL_ENTRY = 1U;
45+constexpr uint32_t OUT_NUM_UNIQUE = 2U;
46+ 
47+constexpr uint32_t ATTR_NUM_ENTRIES_PER_RANK = 0U;
48+constexpr uint32_t ATTR_COMM_BUFFER_SIZE = 1U;
49+ 
50+constexpr uint32_t DIM_ONE = 1U;
51+constexpr uint32_t DIM_TWO = 2U;
52+constexpr uint32_t SYSTEM_NEED_WORKSPACE = 16U * 1024 * 1024;
53+ 
54+constexpr int32_t HIDDEN_SIZE_ALIGN = 128;
55+constexpr int64_t BUFFER_ALIGNMENT = 2 * 1024 * 1024;
56+constexpr int64_t UB_ALIGN = 32;
57+ 
58+static const std::vector<ge::DataType> GRAD_DTYPE_LIST = {ge::DT_BF16, ge::DT_FLOAT16, ge::DT_FLOAT};
59+ 
60+static bool IsContains(const std::vector<ge::DataType> &list, ge::DataType value)
61+{
62+ return std::find(list.begin(), list.end(), value) != list.end();
63+}
64+ 
65+static int64_t CeilDiv(int64_t x, int64_t y)
66+{
67+ return (x + y - 1) / y;
68+}
69+ 
70+static int64_t AlignTo(int64_t x, int64_t y)
71+{
72+ return CeilDiv(x, y) * y;
73+}
74+ 
75+static void PrintEngramFetchGradTilingData(const EngramFetchGradTilingData *tilingData, const char *nodeName)
76+{
77+ OP_TILING_CHECK(tilingData == nullptr, OP_LOGE_WITH_INVALID_INPUT(nodeName, "tilingData"), return);
78+ 
79+ OP_LOGD(nodeName, "========== EngramFetchGradTilingData ==========");
80+ OP_LOGD(nodeName, "numTokens is %ld", tilingData->numTokens);
81+ OP_LOGD(nodeName, "numEntriesPerRank is %d", tilingData->numEntriesPerRank);
82+ OP_LOGD(nodeName, "hiddenDim is %ld", tilingData->hiddenDim);
83+ OP_LOGD(nodeName, "hiddenBytes is %ld", tilingData->hiddenBytes);
84+ OP_LOGD(nodeName, "aivNum is %u", tilingData->aivNum);
85+ OP_LOGD(nodeName, "rankSize is %u", tilingData->rankSize);
86+ OP_LOGD(nodeName, "totalRecv is %ld", tilingData->totalRecv);
87+ OP_LOGD(nodeName, "commBufferSize is %ld", tilingData->commBufferSize);
88+ OP_LOGD(nodeName, "inputDtype is %d", tilingData->inputDtype);
89+ OP_LOGD(nodeName, "outputDtype is %d", tilingData->outputDtype);
90+}
91+ 
92+static ge::graphStatus CheckTensorPtrNullptr(const gert::TilingContext *context)
93+{
94+ auto commContextDesc = context->GetInputDesc(IN_COMM_CONTEXT);
95+ auto gradDesc = context->GetInputDesc(IN_GRAD_FETCHED);
96+ auto permDesc = context->GetInputDesc(IN_PERM);
97+ auto sendCountsDesc = context->GetInputDesc(IN_SEND_COUNTS);
98+ auto recvCountsDesc = context->GetInputDesc(IN_RECV_COUNTS);
99+ auto recvLocalEntryDesc = context->GetInputDesc(IN_RECV_LOCAL_ENTRY);
100+ auto numRecvDesc = context->GetInputDesc(IN_NUM_RECV);
101+ auto gradUniqueDesc = context->GetOutputDesc(OUT_GRAD_UNIQUE);
102+ auto uniqueLocalEntryDesc = context->GetOutputDesc(OUT_UNIQUE_LOCAL_ENTRY);
103+ auto numUniqueDesc = context->GetOutputDesc(OUT_NUM_UNIQUE);
104+ 
105+ OP_CHECK_NULL_WITH_CONTEXT(context, commContextDesc);
106+ OP_CHECK_NULL_WITH_CONTEXT(context, gradDesc);
107+ OP_CHECK_NULL_WITH_CONTEXT(context, permDesc);
108+ OP_CHECK_NULL_WITH_CONTEXT(context, sendCountsDesc);
109+ OP_CHECK_NULL_WITH_CONTEXT(context, recvCountsDesc);
110+ OP_CHECK_NULL_WITH_CONTEXT(context, recvLocalEntryDesc);
111+ OP_CHECK_NULL_WITH_CONTEXT(context, numRecvDesc);
112+ OP_CHECK_NULL_WITH_CONTEXT(context, gradUniqueDesc);
113+ OP_CHECK_NULL_WITH_CONTEXT(context, uniqueLocalEntryDesc);
114+ OP_CHECK_NULL_WITH_CONTEXT(context, numUniqueDesc);
115+ 
116+ return ge::GRAPH_SUCCESS;
117+}
118+ 
119+static ge::graphStatus CheckTensorDataType(const gert::TilingContext *context)
120+{
121+ const char *nodeName = context->GetNodeName();
122+ 
123+ auto commContextDesc = context->GetInputDesc(IN_COMM_CONTEXT);
124+ OP_TILING_CHECK(commContextDesc->GetDataType() != ge::DT_INT32,
125+ OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(nodeName, "commContext",
126+ Ops::Base::ToString(commContextDesc->GetDataType()).c_str(),
127+ "The dtype of commContext must be DT_INT32."),
128+ return ge::GRAPH_FAILED);
129+ 
130+ auto gradDesc = context->GetInputDesc(IN_GRAD_FETCHED);
131+ OP_TILING_CHECK(!IsContains(GRAD_DTYPE_LIST, gradDesc->GetDataType()),
132+ OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(
133+ nodeName, "gradFetched", Ops::Base::ToString(gradDesc->GetDataType()).c_str(),
134+ "The dtype of gradFetched must be DT_BF16, DT_FLOAT16 or DT_FLOAT."),
135+ return ge::GRAPH_FAILED);
136+ 
137+ auto permDesc = context->GetInputDesc(IN_PERM);
138+ OP_TILING_CHECK(permDesc->GetDataType() != ge::DT_INT32,
139+ OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(nodeName, "perm",
140+ Ops::Base::ToString(permDesc->GetDataType()).c_str(),
141+ "The dtype of perm must be DT_INT32."),
142+ return ge::GRAPH_FAILED);
143+ 
144+ auto sendCountsDesc = context->GetInputDesc(IN_SEND_COUNTS);
145+ OP_TILING_CHECK(sendCountsDesc->GetDataType() != ge::DT_INT32,
146+ OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(nodeName, "sendCounts",
147+ Ops::Base::ToString(sendCountsDesc->GetDataType()).c_str(),
148+ "The dtype of sendCounts must be DT_INT32."),
149+ return ge::GRAPH_FAILED);
150+ 
151+ auto recvCountsDesc = context->GetInputDesc(IN_RECV_COUNTS);
152+ OP_TILING_CHECK(recvCountsDesc->GetDataType() != ge::DT_INT32,
153+ OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(nodeName, "recvCounts",
154+ Ops::Base::ToString(recvCountsDesc->GetDataType()).c_str(),
155+ "The dtype of recvCounts must be DT_INT32."),
156+ return ge::GRAPH_FAILED);
157+ 
158+ auto recvLocalEntryDesc = context->GetInputDesc(IN_RECV_LOCAL_ENTRY);
159+ OP_TILING_CHECK(recvLocalEntryDesc->GetDataType() != ge::DT_INT32,
160+ OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(
161+ nodeName, "recvLocalEntry", Ops::Base::ToString(recvLocalEntryDesc->GetDataType()).c_str(),
162+ "The dtype of recvLocalEntry must be DT_INT32."),
163+ return ge::GRAPH_FAILED);
164+ auto numRecvDesc = context->GetInputDesc(IN_NUM_RECV);
165+ OP_TILING_CHECK(numRecvDesc->GetDataType() != ge::DT_INT32,
166+ OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(nodeName, "numRecv",
167+ Ops::Base::ToString(numRecvDesc->GetDataType()).c_str(),
168+ "The dtype of numRecv must be DT_INT32."),
169+ return ge::GRAPH_FAILED);
170+ 
171+ auto gradUniqueDesc = context->GetOutputDesc(OUT_GRAD_UNIQUE);
172+ OP_TILING_CHECK(!IsContains(GRAD_DTYPE_LIST, gradUniqueDesc->GetDataType()),
173+ OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(
174+ nodeName, "gradUniqueOut", Ops::Base::ToString(gradUniqueDesc->GetDataType()).c_str(),
175+ "The dtype of gradUniqueOut must be DT_BF16, DT_FLOAT16 or DT_FLOAT."),
176+ return ge::GRAPH_FAILED);
177+ 
178+ auto uniqueLocalEntryDesc = context->GetOutputDesc(OUT_UNIQUE_LOCAL_ENTRY);
179+ OP_TILING_CHECK(
180+ uniqueLocalEntryDesc->GetDataType() != ge::DT_INT32,
181+ OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(nodeName, "uniqueLocalEntryOut",
182+ Ops::Base::ToString(uniqueLocalEntryDesc->GetDataType()).c_str(),
183+ "The dtype of uniqueLocalEntryOut must be DT_INT32."),
184+ return ge::GRAPH_FAILED);
185+ 
186+ auto numUniqueDesc = context->GetOutputDesc(OUT_NUM_UNIQUE);
187+ OP_TILING_CHECK(numUniqueDesc->GetDataType() != ge::DT_INT32,
188+ OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(nodeName, "numUniqueOut",
189+ Ops::Base::ToString(numUniqueDesc->GetDataType()).c_str(),
190+ "The dtype of numUniqueOut must be DT_INT32."),
191+ return ge::GRAPH_FAILED);
192+ 
193+ return ge::GRAPH_SUCCESS;
194+}
195+ 
196+static ge::graphStatus CheckTensorDim(const gert::TilingContext *context, int64_t &numTokens, uint32_t &rankSize,
197+ int64_t &totalRecv, int64_t &hiddenDim)
198+{
199+ const char *nodeName = context->GetNodeName();
200+ 
201+ const gert::StorageShape *commContextShape = context->GetInputShape(IN_COMM_CONTEXT);
202+ OP_TILING_CHECK(commContextShape == nullptr, OP_LOGE_WITH_INVALID_INPUT(nodeName, "commContext"),
203+ return ge::GRAPH_FAILED);
204+ OP_TILING_CHECK(commContextShape->GetStorageShape().GetDimNum() != DIM_ONE,
205+ OP_LOGE_FOR_INVALID_SHAPEDIM_WITH_REASON(
206+ nodeName, "commContext",
207+ (std::to_string(commContextShape->GetStorageShape().GetDimNum()) + "D").c_str(),
208+ "The shape dim of commContext must be 1D."),
209+ return ge::GRAPH_FAILED);
210+ OP_TILING_CHECK(commContextShape->GetStorageShape().GetDim(0) <= 0,
211+ OP_LOGE_FOR_INVALID_VALUE(
212+ nodeName, "commContext",
213+ (std::string("dim0=") + std::to_string(commContextShape->GetStorageShape().GetDim(0))).c_str(),
214+ "> 0"),
215+ return ge::GRAPH_FAILED);
216+ 
217+ // gradFetched: 2D (T, H)
218+ const gert::StorageShape *gradShape = context->GetInputShape(IN_GRAD_FETCHED);
219+ OP_TILING_CHECK(gradShape == nullptr, OP_LOGE_WITH_INVALID_INPUT(nodeName, "gradFetched"), return ge::GRAPH_FAILED);
220+ OP_TILING_CHECK(gradShape->GetStorageShape().GetDimNum() != DIM_TWO,
221+ OP_LOGE_FOR_INVALID_SHAPEDIM_WITH_REASON(
222+ nodeName, "gradFetched",
223+ (std::to_string(gradShape->GetStorageShape().GetDimNum()) + "D").c_str(),
224+ "The shape dim of gradFetched must be 2D."),
225+ return ge::GRAPH_FAILED);
226+ numTokens = gradShape->GetStorageShape().GetDim(0);
227+ OP_TILING_CHECK(numTokens < 0,
228+ OP_LOGE_FOR_INVALID_VALUE(nodeName, "gradFetched",
229+ (std::string("dim0=") + std::to_string(numTokens)).c_str(), ">= 0"),
230+ return ge::GRAPH_FAILED);
231+ hiddenDim = gradShape->GetStorageShape().GetDim(1);
232+ OP_TILING_CHECK(hiddenDim <= 0,
233+ OP_LOGE_FOR_INVALID_VALUE(nodeName, "gradFetched",
234+ (std::string("dim1=") + std::to_string(hiddenDim)).c_str(), "> 0"),
235+ return ge::GRAPH_FAILED);
236+ OP_TILING_CHECK(
237+ hiddenDim % HIDDEN_SIZE_ALIGN != 0,
238+ OP_LOGE_FOR_INVALID_VALUE(nodeName, "gradFetched", (std::string("dim1=") + std::to_string(hiddenDim)).c_str(),
239+ (std::string("must be ") + std::to_string(HIDDEN_SIZE_ALIGN) + "-aligned").c_str()),
240+ return ge::GRAPH_FAILED);
241+ 
242+ // perm: 1D (T,), dim0 == numTokens
243+ const gert::StorageShape *permShape = context->GetInputShape(IN_PERM);
244+ OP_TILING_CHECK(permShape == nullptr, OP_LOGE_WITH_INVALID_INPUT(nodeName, "perm"), return ge::GRAPH_FAILED);
245+ OP_TILING_CHECK(permShape->GetStorageShape().GetDimNum() != DIM_ONE,
246+ OP_LOGE_FOR_INVALID_SHAPEDIM_WITH_REASON(
247+ nodeName, "perm", (std::to_string(permShape->GetStorageShape().GetDimNum()) + "D").c_str(),
248+ "The shape dim of perm must be 1D."),
249+ return ge::GRAPH_FAILED);
250+ OP_TILING_CHECK(permShape->GetStorageShape().GetDim(0) != numTokens,
251+ OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON(
252+ nodeName, "perm",
253+ (std::string("dim0=") + std::to_string(permShape->GetStorageShape().GetDim(0))).c_str(),
254+ (std::string("dim0 must equal gradFetched dim0=") + std::to_string(numTokens)).c_str()),
255+ return ge::GRAPH_FAILED);
256+ 
257+ // sendCounts: 1D (W,)
258+ const gert::StorageShape *sendCountsShape = context->GetInputShape(IN_SEND_COUNTS);
259+ OP_TILING_CHECK(sendCountsShape == nullptr, OP_LOGE_WITH_INVALID_INPUT(nodeName, "sendCounts"),
260+ return ge::GRAPH_FAILED);
261+ OP_TILING_CHECK(sendCountsShape->GetStorageShape().GetDimNum() != DIM_ONE,
262+ OP_LOGE_FOR_INVALID_SHAPEDIM_WITH_REASON(
263+ nodeName, "sendCounts",
264+ (std::to_string(sendCountsShape->GetStorageShape().GetDimNum()) + "D").c_str(),
265+ "The shape dim of sendCounts must be 1D."),
266+ return ge::GRAPH_FAILED);
267+ OP_TILING_CHECK(sendCountsShape->GetStorageShape().GetDim(0) <= 0,
268+ OP_LOGE_FOR_INVALID_VALUE(
269+ nodeName, "sendCounts",
270+ (std::string("dim0=") + std::to_string(sendCountsShape->GetStorageShape().GetDim(0))).c_str(),
271+ "> 0"),
272+ return ge::GRAPH_FAILED);
273+ rankSize = static_cast<uint32_t>(sendCountsShape->GetStorageShape().GetDim(0));
274+ 
275+ // recvCounts: 1D (W,)
276+ const gert::StorageShape *recvCountsShape = context->GetInputShape(IN_RECV_COUNTS);
277+ OP_TILING_CHECK(recvCountsShape == nullptr, OP_LOGE_WITH_INVALID_INPUT(nodeName, "recvCounts"),
278+ return ge::GRAPH_FAILED);
279+ OP_TILING_CHECK(recvCountsShape->GetStorageShape().GetDimNum() != DIM_ONE,
280+ OP_LOGE_FOR_INVALID_SHAPEDIM_WITH_REASON(
281+ nodeName, "recvCounts",
282+ (std::to_string(recvCountsShape->GetStorageShape().GetDimNum()) + "D").c_str(),
283+ "The shape dim of recvCounts must be 1D."),
284+ return ge::GRAPH_FAILED);
285+ OP_TILING_CHECK(recvCountsShape->GetStorageShape().GetDim(0) <= 0,
286+ OP_LOGE_FOR_INVALID_VALUE(
287+ nodeName, "recvCounts",
288+ (std::string("dim0=") + std::to_string(recvCountsShape->GetStorageShape().GetDim(0))).c_str(),
289+ "> 0"),
290+ return ge::GRAPH_FAILED);
291+ OP_TILING_CHECK(recvCountsShape->GetStorageShape().GetDim(0) != static_cast<int64_t>(rankSize),
292+ OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON(
293+ nodeName, "recvCounts",
294+ (std::string("dim0=") + std::to_string(recvCountsShape->GetStorageShape().GetDim(0))).c_str(),
295+ (std::string("dim0 must equal sendCounts dim0(rankSize)=") + std::to_string(rankSize)).c_str()),
296+ return ge::GRAPH_FAILED);
297+ 
298+ // recvLocalEntry: 1D (R,)
299+ const gert::StorageShape *recvLocalEntryShape = context->GetInputShape(IN_RECV_LOCAL_ENTRY);
300+ OP_TILING_CHECK(recvLocalEntryShape == nullptr, OP_LOGE_WITH_INVALID_INPUT(nodeName, "recvLocalEntry"),
301+ return ge::GRAPH_FAILED);
302+ OP_TILING_CHECK(recvLocalEntryShape->GetStorageShape().GetDimNum() != DIM_ONE,
303+ OP_LOGE_FOR_INVALID_SHAPEDIM_WITH_REASON(
304+ nodeName, "recvLocalEntry",
305+ (std::to_string(recvLocalEntryShape->GetStorageShape().GetDimNum()) + "D").c_str(),
306+ "The shape dim of recvLocalEntry must be 1D."),
307+ return ge::GRAPH_FAILED);
308+ OP_TILING_CHECK(
309+ recvLocalEntryShape->GetStorageShape().GetDim(0) < 0,
310+ OP_LOGE_FOR_INVALID_VALUE(
311+ nodeName, "recvLocalEntry",
312+ (std::string("dim0=") + std::to_string(recvLocalEntryShape->GetStorageShape().GetDim(0))).c_str(), ">= 0"),
313+ return ge::GRAPH_FAILED);
314+ totalRecv = recvLocalEntryShape->GetStorageShape().GetDim(0);
315+ 
316+ // numRecv: 1D, dim0 == 1
317+ const gert::StorageShape *numRecvShape = context->GetInputShape(IN_NUM_RECV);
318+ OP_TILING_CHECK(numRecvShape == nullptr, OP_LOGE_WITH_INVALID_INPUT(nodeName, "numRecv"), return ge::GRAPH_FAILED);
319+ OP_TILING_CHECK(numRecvShape->GetStorageShape().GetDimNum() != DIM_ONE,
320+ OP_LOGE_FOR_INVALID_SHAPEDIM_WITH_REASON(
321+ nodeName, "numRecv",
322+ (std::to_string(numRecvShape->GetStorageShape().GetDimNum()) + "D").c_str(),
323+ "The shape dim of numRecv must be 1D."),
324+ return ge::GRAPH_FAILED);
325+ OP_TILING_CHECK(numRecvShape->GetStorageShape().GetDim(0) != 1,
326+ OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON(
327+ nodeName, "numRecv",
328+ (std::string("dim0=") + std::to_string(numRecvShape->GetStorageShape().GetDim(0))).c_str(),
329+ "dim0 must be 1."),
330+ return ge::GRAPH_FAILED);
331+ 
332+ // gradUniqueOut: 2D
333+ const gert::StorageShape *gradUniqueShape = context->GetOutputShape(OUT_GRAD_UNIQUE);
334+ OP_TILING_CHECK(gradUniqueShape == nullptr, OP_LOGE_WITH_INVALID_INPUT(nodeName, "gradUniqueOut"),
335+ return ge::GRAPH_FAILED);
336+ OP_TILING_CHECK(gradUniqueShape->GetStorageShape().GetDimNum() != DIM_TWO,
337+ OP_LOGE_FOR_INVALID_SHAPEDIM_WITH_REASON(
338+ nodeName, "gradUniqueOut",
339+ (std::to_string(gradUniqueShape->GetStorageShape().GetDimNum()) + "D").c_str(),
340+ "The shape dim of gradUniqueOut must be 2D."),
341+ return ge::GRAPH_FAILED);
342+ 
343+ // uniqueLocalEntryOut: 1D
344+ const gert::StorageShape *uniqueLocalEntryShape = context->GetOutputShape(OUT_UNIQUE_LOCAL_ENTRY);
345+ OP_TILING_CHECK(uniqueLocalEntryShape == nullptr, OP_LOGE_WITH_INVALID_INPUT(nodeName, "uniqueLocalEntryOut"),
346+ return ge::GRAPH_FAILED);
347+ OP_TILING_CHECK(uniqueLocalEntryShape->GetStorageShape().GetDimNum() != DIM_ONE,
348+ OP_LOGE_FOR_INVALID_SHAPEDIM_WITH_REASON(
349+ nodeName, "uniqueLocalEntryOut",
350+ (std::to_string(uniqueLocalEntryShape->GetStorageShape().GetDimNum()) + "D").c_str(),
351+ "The shape dim of uniqueLocalEntryOut must be 1D."),
352+ return ge::GRAPH_FAILED);
353+ 
354+ // numUniqueOut: 1D, dim0 == 1
355+ const gert::StorageShape *numUniqueShape = context->GetOutputShape(OUT_NUM_UNIQUE);
356+ OP_TILING_CHECK(numUniqueShape == nullptr, OP_LOGE_WITH_INVALID_INPUT(nodeName, "numUniqueOut"),
357+ return ge::GRAPH_FAILED);
358+ OP_TILING_CHECK(numUniqueShape->GetStorageShape().GetDimNum() != DIM_ONE,
359+ OP_LOGE_FOR_INVALID_SHAPEDIM_WITH_REASON(
360+ nodeName, "numUniqueOut",
361+ (std::to_string(numUniqueShape->GetStorageShape().GetDimNum()) + "D").c_str(),
362+ "The shape dim of numUniqueOut must be 1D."),
363+ return ge::GRAPH_FAILED);
364+ OP_TILING_CHECK(numUniqueShape->GetStorageShape().GetDim(0) != 1,
365+ OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON(
366+ nodeName, "numUniqueOut",
367+ (std::string("dim0=") + std::to_string(numUniqueShape->GetStorageShape().GetDim(0))).c_str(),
368+ "dim0 must be 1."),
369+ return ge::GRAPH_FAILED);
370+ 
371+ return ge::GRAPH_SUCCESS;
372+}
373+ 
374+static ge::graphStatus CheckTensorFormat(const gert::TilingContext *context)
375+{
376+ const char *nodeName = context->GetNodeName();
377+ 
378+ auto commContextDesc = context->GetInputDesc(IN_COMM_CONTEXT);
379+ ge::Format commContextFormat = static_cast<ge::Format>(ge::GetPrimaryFormat(commContextDesc->GetStorageFormat()));
380+ OP_TILING_CHECK(
381+ commContextFormat != ge::FORMAT_ND,
382+ OP_LOGE_FOR_INVALID_FORMAT(nodeName, "commContext", Ops::Base::ToString(commContextFormat).c_str(), "ND"),
383+ return ge::GRAPH_FAILED);
384+ 
385+ auto gradFetchedDesc = context->GetInputDesc(IN_GRAD_FETCHED);
386+ ge::Format gradFetchedFormat = static_cast<ge::Format>(ge::GetPrimaryFormat(gradFetchedDesc->GetStorageFormat()));
387+ OP_TILING_CHECK(
388+ gradFetchedFormat != ge::FORMAT_ND,
389+ OP_LOGE_FOR_INVALID_FORMAT(nodeName, "gradFetched", Ops::Base::ToString(gradFetchedFormat).c_str(), "ND"),
390+ return ge::GRAPH_FAILED);
391+ 
392+ auto permDesc = context->GetInputDesc(IN_PERM);
393+ ge::Format permFormat = static_cast<ge::Format>(ge::GetPrimaryFormat(permDesc->GetStorageFormat()));
394+ OP_TILING_CHECK(permFormat != ge::FORMAT_ND,
395+ OP_LOGE_FOR_INVALID_FORMAT(nodeName, "perm", Ops::Base::ToString(permFormat).c_str(), "ND"),
396+ return ge::GRAPH_FAILED);
397+ 
398+ auto sendCountsDesc = context->GetInputDesc(IN_SEND_COUNTS);
399+ ge::Format sendCountsFormat = static_cast<ge::Format>(ge::GetPrimaryFormat(sendCountsDesc->GetStorageFormat()));
400+ OP_TILING_CHECK(
401+ sendCountsFormat != ge::FORMAT_ND,
402+ OP_LOGE_FOR_INVALID_FORMAT(nodeName, "sendCounts", Ops::Base::ToString(sendCountsFormat).c_str(), "ND"),
403+ return ge::GRAPH_FAILED);
404+ 
405+ auto recvCountsDesc = context->GetInputDesc(IN_RECV_COUNTS);
406+ ge::Format recvCountsFormat = static_cast<ge::Format>(ge::GetPrimaryFormat(recvCountsDesc->GetStorageFormat()));
407+ OP_TILING_CHECK(
408+ recvCountsFormat != ge::FORMAT_ND,
409+ OP_LOGE_FOR_INVALID_FORMAT(nodeName, "recvCounts", Ops::Base::ToString(recvCountsFormat).c_str(), "ND"),
410+ return ge::GRAPH_FAILED);
411+ 
412+ auto recvLocalEntryDesc = context->GetInputDesc(IN_RECV_LOCAL_ENTRY);
413+ ge::Format recvLocalEntryFormat =
414+ static_cast<ge::Format>(ge::GetPrimaryFormat(recvLocalEntryDesc->GetStorageFormat()));
415+ OP_TILING_CHECK(
416+ recvLocalEntryFormat != ge::FORMAT_ND,
417+ OP_LOGE_FOR_INVALID_FORMAT(nodeName, "recvLocalEntry", Ops::Base::ToString(recvLocalEntryFormat).c_str(), "ND"),
418+ return ge::GRAPH_FAILED);
419+ auto numRecvDesc = context->GetInputDesc(IN_NUM_RECV);
420+ ge::Format numRecvFormat = static_cast<ge::Format>(ge::GetPrimaryFormat(numRecvDesc->GetStorageFormat()));
421+ OP_TILING_CHECK(numRecvFormat != ge::FORMAT_ND,
422+ OP_LOGE_FOR_INVALID_FORMAT(nodeName, "numRecv", Ops::Base::ToString(numRecvFormat).c_str(), "ND"),
423+ return ge::GRAPH_FAILED);
424+ 
425+ auto gradUniqueDesc = context->GetOutputDesc(OUT_GRAD_UNIQUE);
426+ ge::Format gradUniqueFormat = static_cast<ge::Format>(ge::GetPrimaryFormat(gradUniqueDesc->GetStorageFormat()));
427+ OP_TILING_CHECK(
428+ gradUniqueFormat != ge::FORMAT_ND,
429+ OP_LOGE_FOR_INVALID_FORMAT(nodeName, "gradUniqueOut", Ops::Base::ToString(gradUniqueFormat).c_str(), "ND"),
430+ return ge::GRAPH_FAILED);
431+ 
432+ auto uniqueLocalEntryDesc = context->GetOutputDesc(OUT_UNIQUE_LOCAL_ENTRY);
433+ ge::Format uniqueLocalEntryFormat =
434+ static_cast<ge::Format>(ge::GetPrimaryFormat(uniqueLocalEntryDesc->GetStorageFormat()));
435+ OP_TILING_CHECK(uniqueLocalEntryFormat != ge::FORMAT_ND,
436+ OP_LOGE_FOR_INVALID_FORMAT(nodeName, "uniqueLocalEntryOut",
437+ Ops::Base::ToString(uniqueLocalEntryFormat).c_str(), "ND"),
438+ return ge::GRAPH_FAILED);
439+ 
440+ auto numUniqueDesc = context->GetOutputDesc(OUT_NUM_UNIQUE);
441+ ge::Format numUniqueFormat = static_cast<ge::Format>(ge::GetPrimaryFormat(numUniqueDesc->GetStorageFormat()));
442+ OP_TILING_CHECK(
443+ numUniqueFormat != ge::FORMAT_ND,
444+ OP_LOGE_FOR_INVALID_FORMAT(nodeName, "numUniqueOut", Ops::Base::ToString(numUniqueFormat).c_str(), "ND"),
445+ return ge::GRAPH_FAILED);
446+ 
447+ return ge::GRAPH_SUCCESS;
448+}
449+ 
450+static ge::graphStatus TilingCheckEngramFetchGrad(const gert::TilingContext *context, int64_t &numTokens,
451+ uint32_t &rankSize, int64_t &totalRecv, int64_t &hiddenDim)
452+{
453+ const char *nodeName = context->GetNodeName();
454+ 
455+ OP_TILING_CHECK(CheckTensorPtrNullptr(context) != ge::GRAPH_SUCCESS,
456+ OP_LOGE(nodeName, "params check nullptr failed."), return ge::GRAPH_FAILED);
457+ OP_TILING_CHECK(CheckTensorDataType(context) != ge::GRAPH_SUCCESS, OP_LOGE(nodeName, "params dataType is invalid."),
458+ return ge::GRAPH_FAILED);
459+ OP_TILING_CHECK(CheckTensorDim(context, numTokens, rankSize, totalRecv, hiddenDim) != ge::GRAPH_SUCCESS,
460+ OP_LOGE(nodeName, "params shape is invalid."), return ge::GRAPH_FAILED);
461+ OP_TILING_CHECK(CheckTensorFormat(context) != ge::GRAPH_SUCCESS, OP_LOGE(nodeName, "params format is invalid."),
462+ return ge::GRAPH_FAILED);
463+ 
464+ return ge::GRAPH_SUCCESS;
465+}
466+ 
467+static ge::graphStatus CheckAttrs(const gert::TilingContext *context)
468+{
469+ const char *nodeName = context->GetNodeName();
470+ auto attrs = context->GetAttrs();
471+ OP_TILING_CHECK(attrs == nullptr, OP_LOGE_WITH_INVALID_INPUT(nodeName, "attrs"), return ge::GRAPH_FAILED);
472+ 
473+ auto numEntriesPerRankPtr = attrs->GetAttrPointer<int64_t>(ATTR_NUM_ENTRIES_PER_RANK);
474+ OP_TILING_CHECK(numEntriesPerRankPtr == nullptr, OP_LOGE_WITH_INVALID_INPUT(nodeName, "num_entries_per_rank"),
475+ return ge::GRAPH_FAILED);
476+ OP_TILING_CHECK(*numEntriesPerRankPtr < 0 || *numEntriesPerRankPtr > INT32_MAX,
477+ OP_LOGE_FOR_INVALID_VALUE(nodeName, "num_entries_per_rank",
478+ std::to_string(*numEntriesPerRankPtr).c_str(), "[0, INT32_MAX]"),
479+ return ge::GRAPH_FAILED);
480+ 
481+ auto commBufferSizePtr = attrs->GetAttrPointer<int64_t>(ATTR_COMM_BUFFER_SIZE);
482+ OP_TILING_CHECK(commBufferSizePtr == nullptr, OP_LOGE_WITH_INVALID_INPUT(nodeName, "comm_buffer_size"),
483+ return ge::GRAPH_FAILED);
484+ OP_TILING_CHECK(
485+ *commBufferSizePtr <= 0,
486+ OP_LOGE_FOR_INVALID_VALUE(nodeName, "comm_buffer_size", std::to_string(*commBufferSizePtr).c_str(), "> 0"),
487+ return ge::GRAPH_FAILED);
488+ 
489+ return ge::GRAPH_SUCCESS;
490+}
491+ 
492+static ge::graphStatus SetPlatformInfo(gert::TilingContext *context, EngramFetchGradTilingData &tilingData)
493+{
494+ const char *nodeName = context->GetNodeName();
495+ 
496+ auto platformInfo = context->GetPlatformInfo();
497+ OPS_CHECK_NULL_WITH_CONTEXT(context, platformInfo);
498+ auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfo);
499+ uint32_t aivNum = ascendcPlatform.GetCoreNumAiv();
500+ uint32_t aicNum = ascendcPlatform.GetCoreNumAic();
501+ uint32_t numBlocks = ascendcPlatform.CalcTschNumBlocks(aivNum, 0, aivNum);
502+ context->SetBlockDim(numBlocks);
503+ tilingData.aivNum = aivNum;
504+ OP_LOGD(nodeName, "aicNum=%u, aivNum=%u, numBlocks=%u", aicNum, aivNum, numBlocks);
505+ 
506+ return ge::GRAPH_SUCCESS;
507+}
508+ 
509+static ge::graphStatus SetTilingData(gert::TilingContext *context, EngramFetchGradTilingData &tilingData,
510+ int64_t numTokens, uint32_t rankSize, int64_t totalRecv, int64_t hiddenDim)
511+{
512+ const char *nodeName = context->GetNodeName();
513+ auto attrs = context->GetAttrs();
514+ 
515+ tilingData.numTokens = numTokens;
516+ tilingData.rankSize = rankSize;
517+ tilingData.totalRecv = totalRecv;
518+ tilingData.hiddenDim = hiddenDim;
519+ 
520+ auto numEntriesPerRankPtr = attrs->GetAttrPointer<int64_t>(ATTR_NUM_ENTRIES_PER_RANK);
521+ tilingData.numEntriesPerRank = static_cast<int32_t>(*numEntriesPerRankPtr);
522+ 
523+ // hiddenBytes 按输入 dtype(gradFetched)算,作为 a2a 交换 stride
524+ auto gradFetchedDesc = context->GetInputDesc(IN_GRAD_FETCHED);
525+ ge::DataType inputDtype = gradFetchedDesc->GetDataType();
526+ int64_t bytesPerElem = ge::GetSizeByDataType(inputDtype);
527+ OP_TILING_CHECK(tilingData.hiddenDim > INT64_MAX / bytesPerElem,
528+ OP_LOGE(nodeName, "hiddenBytes overflow: hiddenDim=%ld * bytesPerElem=%ld exceeds INT64_MAX",
529+ tilingData.hiddenDim, bytesPerElem),
530+ return ge::GRAPH_FAILED);
531+ tilingData.hiddenBytes = tilingData.hiddenDim * bytesPerElem;
532+ tilingData.inputDtype = static_cast<int32_t>(inputDtype);
533+ 
534+ auto gradUniqueDesc = context->GetOutputDesc(OUT_GRAD_UNIQUE);
535+ tilingData.outputDtype = static_cast<int32_t>(gradUniqueDesc->GetDataType());
536+ 
537+ auto commBufferSizePtr = attrs->GetAttrPointer<int64_t>(ATTR_COMM_BUFFER_SIZE);
538+ tilingData.commBufferSize = *commBufferSizePtr;
539+ 
540+ auto platformInfo = context->GetPlatformInfo();
541+ OPS_CHECK_NULL_WITH_CONTEXT(context, platformInfo);
542+ auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfo);
543+ uint64_t ubSize = 0;
544+ ascendcPlatform.GetCoreMemSize(platform_ascendc::CoreMemType::UB, ubSize);
545+ tilingData.ubSize = ubSize;
546+ 
547+ OP_LOGD(nodeName,
548+ "SetTilingData: numTokens=%ld, hiddenDim=%ld, numEntriesPerRank=%d, hiddenBytes=%ld, "
549+ "ubSize=%lu, rankSize=%u, totalRecv=%ld, commBufferSize=%ld",
550+ tilingData.numTokens, tilingData.hiddenDim, tilingData.numEntriesPerRank, tilingData.hiddenBytes,
551+ tilingData.ubSize, tilingData.rankSize, tilingData.totalRecv, tilingData.commBufferSize);
552+ return ge::GRAPH_SUCCESS;
553+}
554+ 
555+static void SetTilingKey(gert::TilingContext *context)
556+{
557+ const char *nodeName = context->GetNodeName();
558+ const uint64_t tilingKey = GET_TPL_TILING_KEY(ENGRAM_FETCH_GRAD_DEFAULT_MODE);
559+ context->SetTilingKey(tilingKey);
560+ OP_LOGD(nodeName, "tilingKey is [%lu] in engram_fetch_grad.", tilingKey);
561+}
562+ 
563+static ge::graphStatus SetWorkSpace(gert::TilingContext *context, const EngramFetchGradTilingData &tilingData)
564+{
565+ const char *nodeName = context->GetNodeName();
566+ size_t *workSpaces = context->GetWorkspaceSizes(1);
567+ OP_TILING_CHECK(workSpaces == nullptr, OP_LOGE(nodeName, "workSpaces is nullptr."), return ge::GRAPH_FAILED);
568+ 
569+ int64_t numRanks = static_cast<int64_t>(tilingData.rankSize);
570+ int64_t numTokens = tilingData.numTokens;
571+ int64_t hiddenBytes = tilingData.hiddenBytes;
572+ int64_t totalRecv = tilingData.totalRecv;
573+ 
574+ OP_TILING_CHECK(numTokens > 0 && hiddenBytes > INT64_MAX / numTokens,
575+ OP_LOGE(nodeName, "workspace overflow: numTokens=%ld, hiddenBytes=%ld", numTokens, hiddenBytes),
576+ return ge::GRAPH_FAILED);
577+ OP_TILING_CHECK(totalRecv > 0 && hiddenBytes > INT64_MAX / totalRecv,
578+ OP_LOGE(nodeName, "workspace overflow: totalRecv=%ld, hiddenBytes=%ld", totalRecv, hiddenBytes),
579+ return ge::GRAPH_FAILED);
580+ 
581+ int64_t wsGradSorted = numTokens * hiddenBytes;
582+ int64_t wsRecvGrad = totalRecv * hiddenBytes;
583+ int64_t wsSdispls = numRanks * UB_ALIGN;
584+ int64_t wsRdispls = numRanks * UB_ALIGN;
585+ int64_t wsCounterScratch = static_cast<int64_t>(tilingData.aivNum) * UB_ALIGN;
586+ int64_t wsFlagScratch = 32;
587+ 
588+ int64_t wsTotal = wsGradSorted + wsRecvGrad + wsSdispls + wsRdispls + wsCounterScratch + wsFlagScratch;
589+ wsTotal = AlignTo(wsTotal, BUFFER_ALIGNMENT);
590+ wsTotal += SYSTEM_NEED_WORKSPACE;
591+ 
592+ workSpaces[0] = static_cast<size_t>(wsTotal);
593+ OP_LOGD(nodeName,
594+ "backward workspace: gradSorted=%ld, recvGrad=%ld, sdispls=%ld, rdispls=%ld, "
595+ "counterScratch=%ld, flagScratch=%ld, total=%zu",
596+ wsGradSorted, wsRecvGrad, wsSdispls, wsRdispls, wsCounterScratch, wsFlagScratch, workSpaces[0]);
597+ return ge::GRAPH_SUCCESS;
598+}
599+ 
600+static ge::graphStatus EngramFetchGradTilingFunc(gert::TilingContext *context)
601+{
602+ OP_TILING_CHECK(context == nullptr, OP_LOGE("engram_fetch_grad_tiling", "failed to get tiling context."),
603+ return ge::GRAPH_FAILED);
604+ const char *nodeName = context->GetNodeName();
605+ OP_TILING_CHECK(nodeName == nullptr, OP_LOGE_WITH_INVALID_INPUT(nodeName, "nodeName"), return ge::GRAPH_FAILED);
606+ 
607+ OP_LOGI(nodeName, "Enter EngramFetchGrad tiling func.");
608+ 
609+ EngramFetchGradTilingData *tilingData = context->GetTilingData<EngramFetchGradTilingData>();
610+ OP_TILING_CHECK(tilingData == nullptr, OP_LOGE_WITH_INVALID_INPUT(nodeName, "tilingData"), return ge::GRAPH_FAILED);
611+ 
612+ int64_t numTokens = 0;
613+ uint32_t rankSize = 0;
614+ int64_t totalRecv = 0;
615+ int64_t hiddenDim = 0;
616+ OP_TILING_CHECK(TilingCheckEngramFetchGrad(context, numTokens, rankSize, totalRecv, hiddenDim) != ge::GRAPH_SUCCESS,
617+ OP_LOGE(nodeName, "check input/output failed."), return ge::GRAPH_FAILED);
618+ 
619+ OP_TILING_CHECK(CheckAttrs(context) != ge::GRAPH_SUCCESS, OP_LOGE(nodeName, "check attrs failed."),
620+ return ge::GRAPH_FAILED);
621+ 
622+ OP_TILING_CHECK(SetPlatformInfo(context, *tilingData) != ge::GRAPH_SUCCESS,
623+ OP_LOGE(nodeName, "set platform info failed."), return ge::GRAPH_FAILED);
624+ 
625+ OP_TILING_CHECK(SetTilingData(context, *tilingData, numTokens, rankSize, totalRecv, hiddenDim) != ge::GRAPH_SUCCESS,
626+ OP_LOGE(nodeName, "set tiling data failed."), return ge::GRAPH_FAILED);
627+ 
628+ SetTilingKey(context);
629+ 
630+ OP_TILING_CHECK(SetWorkSpace(context, *tilingData) != ge::GRAPH_SUCCESS, OP_LOGE(nodeName, "set workspace failed."),
631+ return ge::GRAPH_FAILED);
632+ 
633+ PrintEngramFetchGradTilingData(tilingData, nodeName);
634+ OP_LOGI(nodeName, "EngramFetchGrad tiling end.");
635+ return ge::GRAPH_SUCCESS;
636+}
637+ 
638+IMPL_OP_OPTILING(EngramFetchGrad).Tiling(EngramFetchGradTilingFunc);
639+} // namespace Mc2Tiling
@@ -0,0 +1,47 @@
1+/**
2+ * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+ * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+ * CANN Open Software License Agreement Version 2.0 (the "License").
5+ * Please refer to the License for details. You may not use this file except in compliance with the License.
6+ * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+ * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+ * See LICENSE in the root of the software repository for the full text of the License.
9+ */
10+ 
11+/*!
12+ * \file engram_fetch_grad.cpp
13+ * \brief EngramFetchGrad 算子 kernel 入口
14+ *
15+ * 单融合 kernel,3 步:
16+ * ① grad_sorted = gradFetched[perm]
17+ * ② a2a(grad_sorted; sendCounts↔recvCounts) → recvGrad (用 commBuffer[peer] = a2a GM)
18+ * ③ owner unique + fp32 scatter-add (recvLocalEntry, recvGrad)
19+ * → gradUniqueOut (fp32) + uniqueLocalEntryOut (int32) + numUniqueOut (int32)
20+ *
21+ * kernel 签名按 op 定义顺序: input0~5, output0~2, workspace, tiling
22+ * input: commContext(0), gradFetched(1), perm(2), sendCounts(3), recvCounts(4), recvLocalEntry(5), numRecv(6)
23+ * output: gradUniqueOut(0), uniqueLocalEntryOut(1), numUniqueOut(2)
24+ */
25+ 
26+#if ASC_DEVKIT_MAJOR >= 9
27+#include "basic_api/kernel_basic_intf.h"
28+#else
29+#include "kernel_operator.h"
30+#endif
31+#include "kernel_tiling/kernel_tiling.h"
32+#include "engram_fetch_grad_tiling_data.h"
33+#include "engram_fetch_grad_tiling_key.h"
34+#include "../../engram_fetch/op_kernel/engram_fetch_utils.h"
35+ 
36+using namespace Mc2Kernel;
37+ 
38+template <uint32_t EngramFetchGradMode>
39+__global__ __aicore__ void engram_fetch_grad(GM_ADDR commContext, GM_ADDR gradFetched, GM_ADDR perm, GM_ADDR sendCounts,
40+ GM_ADDR recvCounts, GM_ADDR recvLocalEntry, GM_ADDR numRecv,
41+ GM_ADDR gradUniqueOut, GM_ADDR uniqueLocalEntryOut, GM_ADDR numUniqueOut,
42+ GM_ADDR workspaceGM, GM_ADDR tilingGM)
43+{
44+ REGISTER_TILING_DEFAULT(EngramFetchGradTilingData);
45+ GET_TILING_DATA_WITH_STRUCT(EngramFetchGradTilingData, tilingData, tilingGM);
46+ KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_AIV_ONLY);
47+}
atomgit-bot
atomgit-botatomgit-bot7月21日

🟡 Medium Priority

两个 kernel 入口函数 engram_fetch_grad(engram_fetch_grad.cpp 第38-55行)和 engram_fetch_train(engram_fetch_train.cpp 第38-57行)均为空桩代码,仅包含 printf("hello ...") 语句,实际计算逻辑(桶排序、a2a 通信、grad 聚合/scatter-add)全部缺失为 // 待实现 注释。

后果:

  • 任何调用这些 kernel 的测试或推理/训练流程将静默成功但输出完全错误(随机未初始化数据或不完整)。
  • UT ascend950_execute_entry(test_aclnn_engram_fetch_grad.cpp 第162-166行)仅检查返回码为"任一合法错误码",不会捕获输出正确性问题。
  • 这些空 kernel 作为功能代码提交,对其他开发者产生误导——看起来接口完整但实际不可用。

建议:完成 kernel 实现后再提交,或如果是有意为之的 WIP 提交,将 kernel 函数体替换为明确的错误返回(如 assert(0) 或设置错误标志),避免静默产生错误结果。

likedislike
不准确?
luozhonglin
luozhonglin
7月24日 评论:
@@ -0,0 +1,34 @@
1+/**
2+ * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+ * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+ * CANN Open Software License Agreement Version 2.0 (the "License").
5+ * Please refer to the License for details. You may not use this file except in compliance with the License.
6+ * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+ * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+ * See LICENSE in the root of the software repository for the full text of the License.
9+ */
10+ 
11+/*!
12+ * \file engram_fetch_grad_tiling_data.h
13+ * \brief kernel侧tiling data结构 — EngramFetchGrad
14+ */
15+ 
16+#ifndef ASCENDC_ENGRAM_FETCH_GRAD_TILING_H
17+#define ASCENDC_ENGRAM_FETCH_GRAD_TILING_H
18+ 
19+#include "kernel_tiling/kernel_tiling.h"
20+ 
21+struct EngramFetchGradTilingData {
22+ int64_t numTokens; // gradFetched / perm dim0
23+ int32_t numEntriesPerRank; // 每 rank entry 数
24+ int64_t hiddenDim; // hidden 维度(从 gradFetched dim1 获取)
25+ int64_t hiddenBytes; // hiddenDim * sizeof(inputDtype)(a2a 交换 stride)
26+ uint32_t aivNum; // AIV 核数
27+ uint64_t ubSize; // UB 空间
28+ uint32_t rankSize; // 通信域 rank 数(从 sendCounts dim0 获取)
29+ int64_t totalRecv; // recvLocalEntry dim0(a2a 接收上界)
30+ int64_t commBufferSize; // a2a GM 收发缓冲大小(即 commBuffer 200MB)
31+ int32_t inputDtype; // gradFetched 的 dtype(ge::DataType)
32+ int32_t outputDtype; // gradUniqueOut 的 dtype(ge::DataType)
33+};
34+#endif
@@ -0,0 +1,29 @@
1+/**
2+ * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+ * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+ * CANN Open Software License Agreement Version 2.0 (the "License").
5+ * Please refer to the License for details. You may not use this file except in compliance with the License.
6+ * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+ * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+ * See LICENSE in the root of the software repository for the full text of the License.
9+ */
10+ 
11+/*!
12+ * \file engram_fetch_grad_tiling_key.h
13+ * \brief kernel侧tiling key
14+ */
15+ 
16+#ifndef ENGRAM_FETCH_GRAD_TILING_KEY_H
17+#define ENGRAM_FETCH_GRAD_TILING_KEY_H
18+ 
19+#include "ascendc/host_api/tiling/template_argument.h"
20+ 
21+#define ENGRAM_FETCH_GRAD_DEFAULT_MODE 0
22+ 
23+ASCENDC_TPL_ARGS_DECL(EngramFetchGrad, ASCENDC_TPL_UINT_DECL(EngramFetchGradMode, ASCENDC_TPL_8_BW, ASCENDC_TPL_UI_LIST,
24+ ENGRAM_FETCH_GRAD_DEFAULT_MODE));
25+ 
26+ASCENDC_TPL_SEL(ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(EngramFetchGradMode, ASCENDC_TPL_UI_LIST,
27+ ENGRAM_FETCH_GRAD_DEFAULT_MODE)));
28+ 
29+#endif // ENGRAM_FETCH_GRAD_TILING_KEY_H
@@ -0,0 +1,16 @@
1+# -----------------------------------------------------------------------------------------------------------
2+# Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+# This program is free software; you can redistribute it and/or modify it under the terms and conditions of
4+# CANN Open Software License Agreement Version 2.0 (the "License").
5+# Please refer to the License for details. You may not use this file except in compliance with the License.
6+# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+# See LICENSE in the root of the software repository for the full text of the License.
9+# -----------------------------------------------------------------------------------------------------------
10+ 
11+file(GLOB CURRENT_DIRS RELATIVE ${CMAKE_CURRENT_SOURCE_DIR} ${CMAKE_CURRENT_SOURCE_DIR}/*)
12+foreach(SUB_DIR ${CURRENT_DIRS})
13+ if(EXISTS "${CMAKE_CURRENT_SOURCE_DIR}/${SUB_DIR}/CMakeLists.txt")
14+ add_subdirectory(${SUB_DIR})
15+ endif()
16+endforeach()
@@ -0,0 +1,16 @@
1+# -----------------------------------------------------------------------------------------------------------
2+# Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+# This program is free software; you can redistribute it and/or modify it under the terms and conditions of
4+# CANN Open Software License Agreement Version 2.0 (the "License").
5+# Please refer to the License for details. You may not use this file except in compliance with the License.
6+# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+# See LICENSE in the root of the software repository for the full text of the License.
9+# -----------------------------------------------------------------------------------------------------------
10+ 
11+file(GLOB CURRENT_DIRS RELATIVE ${CMAKE_CURRENT_SOURCE_DIR} ${CMAKE_CURRENT_SOURCE_DIR}/*)
12+foreach(SUB_DIR ${CURRENT_DIRS})
13+ if(EXISTS "${CMAKE_CURRENT_SOURCE_DIR}/${SUB_DIR}/CMakeLists.txt")
14+ add_subdirectory(${SUB_DIR})
15+ endif()
16+endforeach()
@@ -0,0 +1,13 @@
1+# -----------------------------------------------------------------------------------------------------------
2+# Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+# This program is free software; you can redistribute it and/or modify it under the terms and conditions of
4+# CANN Open Software License Agreement Version 2.0 (the "License").
5+# Please refer to the License for details. You may not use this file except in compliance with the License.
6+# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+# See LICENSE in the root of the software repository for the full text of the License.
9+# -----------------------------------------------------------------------------------------------------------
10+ 
11+if(UT_TEST_ALL OR OP_API_UT)
12+ add_modules_ut_sources(UT_NAME ${OP_API_MODULE_NAME} MODE PRIVATE DIR ${CMAKE_CURRENT_SOURCE_DIR})
13+endif()
@@ -0,0 +1,177 @@
1+/**
2+ * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+ * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+ * CANN Open Software License Agreement Version 2.0 (the "License").
5+ * Please refer to the License for details. You may not use this file except in compliance with the License.
6+ * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+ * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+ * See LICENSE in the root of the software repository for the full text of the License.
9+ */
10+ 
11+/*!
12+ * \file test_aclnn_engram_fetch_grad.cpp
13+ * \brief EngramFetchGrad 算子 op_api 侧 aclnn 接口 UT
14+ *
15+ * 测试 aclnnEngramFetchGradGetWorkspaceSize 的参数校验:
16+ * - 正常场景: 合法输入输出
17+ * - nullptr 场景: 各输入/输出为 nullptr
18+ *
19+ * aclnnEngramFetchGrad 接口签名:
20+ * input: commContext, gradFetched, perm, sendCounts, recvCounts, recvLocalEntry, numRecv
21+ * output: gradUniqueOut, uniqueLocalEntryOut, numUniqueOut
22+ * attr: numEntriesPerRank, commBufferSize
23+ */
24+ 
25+#include <gtest/gtest.h>
26+#include <gmock/gmock.h>
27+#include "op_api_ut_common/op_api_ut.h"
28+#include "op_api_ut_common/tensor_desc.h"
29+#include "opdev/platform.h"
30+#include "aclnn/aclnn_base.h"
31+ 
32+extern "C" {
33+aclnnStatus aclnnEngramFetchGradGetWorkspaceSize(const aclTensor *commContext, const aclTensor *gradFetched,
34+ const aclTensor *perm, const aclTensor *sendCounts,
35+ const aclTensor *recvCounts, const aclTensor *recvLocalEntry,
36+ const aclTensor *numRecv, aclTensor *gradUniqueOut,
37+ aclTensor *uniqueLocalEntryOut, aclTensor *numUniqueOut,
38+ int64_t numEntriesPerRank, int64_t commBufferSize,
39+ uint64_t *workspaceSize, aclOpExecutor **executor);
40+ 
41+aclnnStatus aclnnEngramFetchGrad(void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, aclrtStream stream);
42+}
atomgit-bot
atomgit-botatomgit-bot7月21日

🔴 Critical

测试文件 test_aclnn_engram_fetch_grad.cpp 第 33-39 行的 extern "C" 声明与 aclnn_engram_fetch_grad.cpp 第 63-70 行的实际实现签名存在严重不匹配:

这导致测试中的 INPUT(commContext_desc, gradFetched_desc, perm_desc, sendCounts_desc, recvCounts_desc, recvLocalEntry_desc, hiddenSize, numEntriesPerRank) 与实际函数签名不匹配——参数类型、数量和顺序全部错位,会引发 ODR 违规、未定义行为,测试必然崩溃或产生错误结果。同时测试用例缺少 numRecv tensor(实现要求 7 个输入),而测试仅提供了 6 个 tensor。

建议:将 extern 声明改为与实际实现签名一致:在 recvLocalEntry 之后增加 const aclTensor *numRecv 参数,在 3 个输出 tensor 之后使用 int64_t numEntriesPerRank, int64_t commBufferSize 作为标量参数;同时所有测试用例需增加 numRecv_desc tensor 输入并移除不存在的 hiddenSize 标量参数。

likedislike
不准确?
luozhonglin
luozhonglin
7月24日 评论:
43+ 
44+using namespace op;
45+using namespace std;
46+ 
47+class AclnnEngramFetchGradTest : public testing::Test {
48+protected:
49+ static void SetUpTestCase()
50+ {
51+ op::SetPlatformSocVersion(op::SocVersion::ASCEND950);
52+ std::cout << "EngramFetchGrad AclnnTest SetUp" << std::endl;
53+ }
54+ 
55+ static void TearDownTestCase()
56+ {
57+ op::SetPlatformSocVersion(op::SocVersion::ASCEND950);
58+ std::cout << "EngramFetchGrad AclnnTest TearDown" << std::endl;
59+ }
60+};
61+ 
62+TEST_F(AclnnEngramFetchGradTest, ascend950_success)
63+{
64+ auto commContext_desc = TensorDesc({6146}, ACL_INT32, ACL_FORMAT_ND).ValueRange(0, 1);
65+ auto gradFetched_desc = TensorDesc({8, 512}, ACL_BF16, ACL_FORMAT_ND);
66+ auto perm_desc = TensorDesc({8}, ACL_INT32, ACL_FORMAT_ND).ValueRange(0, 7);
67+ auto sendCounts_desc = TensorDesc({2}, ACL_INT32, ACL_FORMAT_ND).ValueRange(0, 8);
68+ auto recvCounts_desc = TensorDesc({2}, ACL_INT32, ACL_FORMAT_ND).ValueRange(0, 16);
69+ auto recvLocalEntry_desc = TensorDesc({16}, ACL_INT32, ACL_FORMAT_ND).ValueRange(0, 3);
70+ auto numRecv_desc = TensorDesc({1}, ACL_INT32, ACL_FORMAT_ND);
71+ auto gradUnique_desc = TensorDesc({16, 512}, ACL_BF16, ACL_FORMAT_ND);
72+ auto uniqueLocalEntry_desc = TensorDesc({16}, ACL_INT32, ACL_FORMAT_ND);
73+ auto numUnique_desc = TensorDesc({1}, ACL_INT32, ACL_FORMAT_ND);
74+ 
75+ int64_t numEntriesPerRank = 4;
76+ int64_t commBufferSize = 4194304;
77+ 
78+ auto ut = OP_API_UT(aclnnEngramFetchGrad,
79+ INPUT(commContext_desc, gradFetched_desc, perm_desc, sendCounts_desc, recvCounts_desc,
80+ recvLocalEntry_desc, numRecv_desc, gradUnique_desc, uniqueLocalEntry_desc, numUnique_desc,
81+ numEntriesPerRank, commBufferSize),
82+ OUTPUT());
83+ 
84+ uint64_t workspace_size = 0;
85+ aclOpExecutor *executor = nullptr;
86+ aclnnStatus aclRet = ut.TestGetWorkspaceSizeWithNNopbaseInner(&workspace_size, executor);
87+ EXPECT_NE(aclRet, ACLNN_ERR_PARAM_NULLPTR);
88+ EXPECT_NE(aclRet, ACLNN_ERR_PARAM_INVALID);
89+}
90+ 
91+TEST_F(AclnnEngramFetchGradTest, ascend950_nullptr_commContext)
92+{
93+ auto gradFetched_desc = TensorDesc({8, 512}, ACL_BF16, ACL_FORMAT_ND);
94+ auto perm_desc = TensorDesc({8}, ACL_INT32, ACL_FORMAT_ND);
95+ auto sendCounts_desc = TensorDesc({2}, ACL_INT32, ACL_FORMAT_ND);
96+ auto recvCounts_desc = TensorDesc({2}, ACL_INT32, ACL_FORMAT_ND);
97+ auto recvLocalEntry_desc = TensorDesc({16}, ACL_INT32, ACL_FORMAT_ND);
98+ auto numRecv_desc = TensorDesc({1}, ACL_INT32, ACL_FORMAT_ND);
99+ auto gradUnique_desc = TensorDesc({16, 512}, ACL_BF16, ACL_FORMAT_ND);
100+ auto uniqueLocalEntry_desc = TensorDesc({16}, ACL_INT32, ACL_FORMAT_ND);
101+ auto numUnique_desc = TensorDesc({1}, ACL_INT32, ACL_FORMAT_ND);
102+ 
103+ int64_t numEntriesPerRank = 4;
104+ int64_t commBufferSize = 4194304;
105+ 
106+ auto ut = OP_API_UT(aclnnEngramFetchGrad,
107+ INPUT(nullptr, gradFetched_desc, perm_desc, sendCounts_desc, recvCounts_desc,
108+ recvLocalEntry_desc, numRecv_desc, gradUnique_desc, uniqueLocalEntry_desc, numUnique_desc,
109+ numEntriesPerRank, commBufferSize),
110+ OUTPUT());
111+ 
112+ uint64_t workspace_size = 0;
113+ aclOpExecutor *executor = nullptr;
114+ aclnnStatus aclRet = ut.TestGetWorkspaceSizeWithNNopbaseInner(&workspace_size, executor);
115+ EXPECT_EQ(aclRet, ACLNN_ERR_PARAM_NULLPTR);
116+}
117+ 
118+TEST_F(AclnnEngramFetchGradTest, ascend950_nullptr_gradFetched)
119+{
120+ auto commContext_desc = TensorDesc({6146}, ACL_INT32, ACL_FORMAT_ND);
121+ auto perm_desc = TensorDesc({8}, ACL_INT32, ACL_FORMAT_ND);
122+ auto sendCounts_desc = TensorDesc({2}, ACL_INT32, ACL_FORMAT_ND);
123+ auto recvCounts_desc = TensorDesc({2}, ACL_INT32, ACL_FORMAT_ND);
124+ auto recvLocalEntry_desc = TensorDesc({16}, ACL_INT32, ACL_FORMAT_ND);
125+ auto numRecv_desc = TensorDesc({1}, ACL_INT32, ACL_FORMAT_ND);
126+ auto gradUnique_desc = TensorDesc({16, 512}, ACL_BF16, ACL_FORMAT_ND);
127+ auto uniqueLocalEntry_desc = TensorDesc({16}, ACL_INT32, ACL_FORMAT_ND);
128+ auto numUnique_desc = TensorDesc({1}, ACL_INT32, ACL_FORMAT_ND);
129+ 
130+ int64_t numEntriesPerRank = 4;
131+ int64_t commBufferSize = 4194304;
132+ 
133+ auto ut = OP_API_UT(aclnnEngramFetchGrad,
134+ INPUT(commContext_desc, nullptr, perm_desc, sendCounts_desc, recvCounts_desc,
135+ recvLocalEntry_desc, numRecv_desc, gradUnique_desc, uniqueLocalEntry_desc, numUnique_desc,
136+ numEntriesPerRank, commBufferSize),
137+ OUTPUT());
138+ 
139+ uint64_t workspace_size = 0;
140+ aclOpExecutor *executor = nullptr;
141+ aclnnStatus aclRet = ut.TestGetWorkspaceSizeWithNNopbaseInner(&workspace_size, executor);
142+ EXPECT_EQ(aclRet, ACLNN_ERR_PARAM_NULLPTR);
143+}
144+ 
145+TEST_F(AclnnEngramFetchGradTest, ascend950_nullptr_gradUniqueOut)
146+{
147+ auto commContext_desc = TensorDesc({6146}, ACL_INT32, ACL_FORMAT_ND);
148+ auto gradFetched_desc = TensorDesc({8, 512}, ACL_BF16, ACL_FORMAT_ND);
149+ auto perm_desc = TensorDesc({8}, ACL_INT32, ACL_FORMAT_ND);
150+ auto sendCounts_desc = TensorDesc({2}, ACL_INT32, ACL_FORMAT_ND);
151+ auto recvCounts_desc = TensorDesc({2}, ACL_INT32, ACL_FORMAT_ND);
152+ auto recvLocalEntry_desc = TensorDesc({16}, ACL_INT32, ACL_FORMAT_ND);
153+ auto numRecv_desc = TensorDesc({1}, ACL_INT32, ACL_FORMAT_ND);
154+ auto uniqueLocalEntry_desc = TensorDesc({16}, ACL_INT32, ACL_FORMAT_ND);
155+ auto numUnique_desc = TensorDesc({1}, ACL_INT32, ACL_FORMAT_ND);
156+ 
157+ int64_t numEntriesPerRank = 4;
158+ int64_t commBufferSize = 4194304;
159+ 
160+ auto ut = OP_API_UT(aclnnEngramFetchGrad,
161+ INPUT(commContext_desc, gradFetched_desc, perm_desc, sendCounts_desc, recvCounts_desc,
162+ recvLocalEntry_desc, numRecv_desc, nullptr, uniqueLocalEntry_desc, numUnique_desc,
163+ numEntriesPerRank, commBufferSize),
164+ OUTPUT());
165+ 
166+ uint64_t workspace_size = 0;
167+ aclOpExecutor *executor = nullptr;
168+ aclnnStatus aclRet = ut.TestGetWorkspaceSizeWithNNopbaseInner(&workspace_size, executor);
169+ EXPECT_EQ(aclRet, ACLNN_ERR_PARAM_NULLPTR);
170+}
171+ 
172+TEST_F(AclnnEngramFetchGradTest, ascend950_execute_entry)
173+{
174+ aclnnStatus ret = aclnnEngramFetchGrad(nullptr, 0, nullptr, nullptr);
175+ EXPECT_THAT(ret, testing::AnyOf(testing::Eq(ACLNN_SUCCESS), testing::Eq(ACLNN_ERR_PARAM_NULLPTR),
176+ testing::Eq(ACLNN_ERR_PARAM_INVALID), testing::Eq(ACLNN_ERR_INNER)));
177+}
@@ -0,0 +1,16 @@
1+# -----------------------------------------------------------------------------------------------------------
2+# Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+# This program is free software; you can redistribute it and/or modify it under the terms and conditions of
4+# CANN Open Software License Agreement Version 2.0 (the "License").
5+# Please refer to the License for details. You may not use this file except in compliance with the License.
6+# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+# See LICENSE in the root of the software repository for the full text of the License.
9+# -----------------------------------------------------------------------------------------------------------
10+ 
11+file(GLOB CURRENT_DIRS RELATIVE ${CMAKE_CURRENT_SOURCE_DIR} ${CMAKE_CURRENT_SOURCE_DIR}/*)
12+foreach(SUB_DIR ${CURRENT_DIRS})
13+ if(EXISTS "${CMAKE_CURRENT_SOURCE_DIR}/${SUB_DIR}/CMakeLists.txt")
14+ add_subdirectory(${SUB_DIR})
15+ endif()
16+endforeach()
@@ -0,0 +1,13 @@
1+# -----------------------------------------------------------------------------------------------------------
2+# Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+# This program is free software; you can redistribute it and/or modify it under the terms and conditions of
4+# CANN Open Software License Agreement Version 2.0 (the "License").
5+# Please refer to the License for details. You may not use this file except in compliance with the License.
6+# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+# See LICENSE in the root of the software repository for the full text of the License.
9+# -----------------------------------------------------------------------------------------------------------
10+ 
11+if(UT_TEST_ALL OR OP_HOST_UT)
12+ add_modules_ut_sources(UT_NAME ${OP_TILING_MODULE_NAME} MODE PRIVATE DIR ${CMAKE_CURRENT_SOURCE_DIR})
13+endif()
@@ -0,0 +1,489 @@
1+/**
2+ * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+ * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+ * CANN Open Software License Agreement Version 2.0 (the "License").
5+ * Please refer to the License for details. You may not use this file except in compliance with the License.
6+ * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+ * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+ * See LICENSE in the root of the software repository for the full text of the License.
9+ */
10+ 
11+/*!
12+ * \file test_engram_fetch_grad_tiling.cpp
13+ * \brief EngramFetchGrad 算子 host 侧 tiling UT
14+ *
15+ * 测试用例覆盖:
16+ * - 正常场景: bf16/fp16/fp32 gradFetched, 各种 token/hidden 组合
17+ * - 异常场景: dtype 不匹配、维度错误、attr 非法值
18+ *
19+ * EngramFetchGrad 输入输出规格:
20+ * input: commContext(0), gradFetched(1), perm(2), sendCounts(3), recvCounts(4), recvLocalEntry(5), numRecv(6)
21+ * output: gradUniqueOut(0), uniqueLocalEntryOut(1), numUniqueOut(2)
22+ * attr: num_entries_per_rank, comm_buffer_size
23+ */
24+ 
25+#include <iostream>
26+#include <gtest/gtest.h>
27+#include "tiling_case_executor.h"
28+ 
29+namespace EngramFetchGradUT {
30+ 
31+static const std::string OP_NAME = "EngramFetchGrad";
32+ 
33+struct EngramFetchGradTestParam {
34+ std::string caseName;
35+ std::initializer_list<int64_t> commContextShape;
36+ ge::DataType commContextDtype;
37+ ge::Format commContextFormat;
38+ std::initializer_list<int64_t> gradFetchedShape;
39+ ge::DataType gradFetchedDtype;
40+ ge::Format gradFetchedFormat;
41+ std::initializer_list<int64_t> permShape;
42+ ge::DataType permDtype;
43+ std::initializer_list<int64_t> sendCountsShape;
44+ std::initializer_list<int64_t> recvCountsShape;
45+ std::initializer_list<int64_t> recvLocalEntryShape;
46+ std::initializer_list<int64_t> gradUniqueShape;
47+ ge::DataType gradUniqueDtype;
48+ std::initializer_list<int64_t> uniqueLocalEntryShape;
49+ std::initializer_list<int64_t> numUniqueShape;
50+ int64_t numEntriesPerRank;
51+ int64_t commBufferSize;
52+ std::string socVersion;
53+ ge::graphStatus status;
54+ uint64_t expectTilingKey;
55+ std::string expectTilingData;
56+ std::vector<size_t> expectWorkspaces;
57+};
58+ 
59+inline std::ostream &operator<<(std::ostream &os, const EngramFetchGradTestParam &param)
60+{
61+ return os << param.caseName;
62+}
63+ 
64+static EngramFetchGradTestParam g_testCases[] = {
65+ {"success_bf16_basic",
66+ {6146},
67+ ge::DT_INT32,
68+ ge::FORMAT_ND,
69+ {8, 512},
70+ ge::DT_BF16,
71+ ge::FORMAT_ND,
72+ {8},
73+ ge::DT_INT32,
74+ {2},
75+ {2},
76+ {16},
77+ {16, 512},
78+ ge::DT_BF16,
79+ {16},
80+ {1},
81+ 4,
82+ 4194304,
83+ "3510",
84+ ge::GRAPH_SUCCESS,
85+ 0UL,
86+ "",
87+ {}},
88+ 
89+ {"success_fp16_basic",
90+ {6146},
91+ ge::DT_INT32,
92+ ge::FORMAT_ND,
93+ {8, 512},
94+ ge::DT_FLOAT16,
95+ ge::FORMAT_ND,
96+ {8},
97+ ge::DT_INT32,
98+ {2},
99+ {2},
100+ {16},
101+ {16, 512},
102+ ge::DT_FLOAT16,
103+ {16},
104+ {1},
105+ 4,
106+ 4194304,
107+ "3510",
108+ ge::GRAPH_SUCCESS,
109+ 0UL,
110+ "",
111+ {}},
112+ 
113+ {"success_fp32_basic",
114+ {6146},
115+ ge::DT_INT32,
116+ ge::FORMAT_ND,
117+ {8, 512},
118+ ge::DT_FLOAT,
119+ ge::FORMAT_ND,
120+ {8},
121+ ge::DT_INT32,
122+ {2},
123+ {2},
124+ {16},
125+ {16, 512},
126+ ge::DT_FLOAT,
127+ {16},
128+ {1},
129+ 4,
130+ 4194304,
131+ "3510",
132+ ge::GRAPH_SUCCESS,
133+ 0UL,
134+ "",
135+ {}},
136+ 
137+ {"success_large_tokens",
138+ {6146},
139+ ge::DT_INT32,
140+ ge::FORMAT_ND,
141+ {128, 256},
142+ ge::DT_BF16,
143+ ge::FORMAT_ND,
144+ {128},
145+ ge::DT_INT32,
146+ {2},
147+ {2},
148+ {256},
149+ {256, 256},
150+ ge::DT_BF16,
151+ {256},
152+ {1},
153+ 100,
154+ 4194304,
155+ "3510",
156+ ge::GRAPH_SUCCESS,
157+ 0UL,
158+ "",
159+ {}},
160+ 
161+ {"success_single_token",
162+ {6146},
163+ ge::DT_INT32,
164+ ge::FORMAT_ND,
165+ {1, 512},
166+ ge::DT_BF16,
167+ ge::FORMAT_ND,
168+ {1},
169+ ge::DT_INT32,
170+ {2},
171+ {2},
172+ {2},
173+ {2, 512},
174+ ge::DT_BF16,
175+ {2},
176+ {1},
177+ 4,
178+ 4194304,
179+ "3510",
180+ ge::GRAPH_SUCCESS,
181+ 0UL,
182+ "",
183+ {}},
184+ 
185+ {"fail_gradFetched_dtype_int32",
186+ {6146},
187+ ge::DT_INT32,
188+ ge::FORMAT_ND,
189+ {8, 512},
190+ ge::DT_INT32,
191+ ge::FORMAT_ND,
192+ {8},
193+ ge::DT_INT32,
194+ {2},
195+ {2},
196+ {16},
197+ {16, 512},
198+ ge::DT_INT32,
199+ {16},
200+ {1},
201+ 4,
202+ 4194304,
203+ "3510",
204+ ge::GRAPH_FAILED,
205+ 0UL,
206+ "",
207+ {}},
208+ 
209+ {"fail_gradUnique_dtype_int32",
210+ {6146},
211+ ge::DT_INT32,
212+ ge::FORMAT_ND,
213+ {8, 512},
214+ ge::DT_BF16,
215+ ge::FORMAT_ND,
216+ {8},
217+ ge::DT_INT32,
218+ {2},
219+ {2},
220+ {16},
221+ {16, 512},
222+ ge::DT_INT32,
223+ {16},
224+ {1},
225+ 4,
226+ 4194304,
227+ "3510",
228+ ge::GRAPH_FAILED,
229+ 0UL,
230+ "",
231+ {}},
232+ 
233+ {"fail_perm_dtype_float",
234+ {6146},
235+ ge::DT_INT32,
236+ ge::FORMAT_ND,
237+ {8, 512},
238+ ge::DT_BF16,
239+ ge::FORMAT_ND,
240+ {8},
241+ ge::DT_FLOAT,
242+ {2},
243+ {2},
244+ {16},
245+ {16, 512},
246+ ge::DT_BF16,
247+ {16},
248+ {1},
249+ 4,
250+ 4194304,
251+ "3510",
252+ ge::GRAPH_FAILED,
253+ 0UL,
254+ "",
255+ {}},
256+ 
257+ {"fail_gradFetched_1d",
258+ {6146},
259+ ge::DT_INT32,
260+ ge::FORMAT_ND,
261+ {4096},
262+ ge::DT_BF16,
263+ ge::FORMAT_ND,
264+ {8},
265+ ge::DT_INT32,
266+ {2},
267+ {2},
268+ {16},
269+ {16, 512},
270+ ge::DT_BF16,
271+ {16},
272+ {1},
273+ 4,
274+ 4194304,
275+ "3510",
276+ ge::GRAPH_FAILED,
277+ 0UL,
278+ "",
279+ {}},
280+ 
281+ {"fail_perm_dim0_mismatch",
282+ {6146},
283+ ge::DT_INT32,
284+ ge::FORMAT_ND,
285+ {8, 512},
286+ ge::DT_BF16,
287+ ge::FORMAT_ND,
288+ {16},
289+ ge::DT_INT32,
290+ {2},
291+ {2},
292+ {16},
293+ {16, 512},
294+ ge::DT_BF16,
295+ {16},
296+ {1},
297+ 4,
298+ 4194304,
299+ "3510",
300+ ge::GRAPH_FAILED,
301+ 0UL,
302+ "",
303+ {}},
304+ 
305+ {"fail_numUnique_dim0_not_1",
306+ {6146},
307+ ge::DT_INT32,
308+ ge::FORMAT_ND,
309+ {8, 512},
310+ ge::DT_BF16,
311+ ge::FORMAT_ND,
312+ {8},
313+ ge::DT_INT32,
314+ {2},
315+ {2},
316+ {16},
317+ {16, 512},
318+ ge::DT_BF16,
319+ {16},
320+ {2},
321+ 4,
322+ 4194304,
323+ "3510",
324+ ge::GRAPH_FAILED,
325+ 0UL,
326+ "",
327+ {}},
328+ 
329+ {"fail_hidden_size_zero",
330+ {6146},
331+ ge::DT_INT32,
332+ ge::FORMAT_ND,
333+ {8, 0},
334+ ge::DT_BF16,
335+ ge::FORMAT_ND,
336+ {8},
337+ ge::DT_INT32,
338+ {2},
339+ {2},
340+ {16},
341+ {16, 0},
342+ ge::DT_BF16,
343+ {16},
344+ {1},
345+ 4,
346+ 4194304,
347+ "3510",
348+ ge::GRAPH_FAILED,
349+ 0UL,
350+ "",
351+ {}},
352+ 
353+ {"fail_hidden_size_not_aligned",
354+ {6146},
355+ ge::DT_INT32,
356+ ge::FORMAT_ND,
357+ {8, 100},
358+ ge::DT_BF16,
359+ ge::FORMAT_ND,
360+ {8},
361+ ge::DT_INT32,
362+ {2},
363+ {2},
364+ {16},
365+ {16, 100},
366+ ge::DT_BF16,
367+ {16},
368+ {1},
369+ 4,
370+ 4194304,
371+ "3510",
372+ ge::GRAPH_FAILED,
373+ 0UL,
374+ "",
375+ {}},
376+ 
377+ {"fail_num_entries_negative",
378+ {6146},
379+ ge::DT_INT32,
380+ ge::FORMAT_ND,
381+ {8, 512},
382+ ge::DT_BF16,
383+ ge::FORMAT_ND,
384+ {8},
385+ ge::DT_INT32,
386+ {2},
387+ {2},
388+ {16},
389+ {16, 512},
390+ ge::DT_BF16,
391+ {16},
392+ {1},
393+ -1,
394+ 4194304,
395+ "3510",
396+ ge::GRAPH_FAILED,
397+ 0UL,
398+ "",
399+ {}},
400+ 
401+ {"fail_commContext_dtype_bf16",
402+ {6146},
403+ ge::DT_BF16,
404+ ge::FORMAT_ND,
405+ {8, 512},
406+ ge::DT_BF16,
407+ ge::FORMAT_ND,
408+ {8},
409+ ge::DT_INT32,
410+ {2},
411+ {2},
412+ {16},
413+ {16, 512},
414+ ge::DT_BF16,
415+ {16},
416+ {1},
417+ 4,
418+ 4194304,
419+ "3510",
420+ ge::GRAPH_FAILED,
421+ 0UL,
422+ "",
423+ {}},
424+};
425+ 
426+class EngramFetchGradArch35TilingTest : public testing::TestWithParam<EngramFetchGradTestParam> {
427+protected:
428+ static void SetUpTestCase()
429+ {
430+ std::cout << "EngramFetchGradArch35TilingTest SetUp." << std::endl;
431+ }
432+ 
433+ static void TearDownTestCase()
434+ {
435+ std::cout << "EngramFetchGradArch35TilingTest TearDown." << std::endl;
436+ }
437+};
438+ 
439+static struct EngramFetchGradCompileInfo {
440+} compileInfo;
441+ 
442+static gert::TilingContextPara BuildTilingContextPara(const EngramFetchGradTestParam &param)
443+{
444+ std::cout << "[TEST_CASE] " << param.caseName << std::endl;
445+ gert::StorageShape commContextShape = {param.commContextShape, param.commContextShape};
446+ gert::StorageShape gradFetchedShape = {param.gradFetchedShape, param.gradFetchedShape};
447+ gert::StorageShape permShape = {param.permShape, param.permShape};
448+ gert::StorageShape sendCountsShape = {param.sendCountsShape, param.sendCountsShape};
449+ gert::StorageShape recvCountsShape = {param.recvCountsShape, param.recvCountsShape};
450+ gert::StorageShape recvLocalEntryShape = {param.recvLocalEntryShape, param.recvLocalEntryShape};
451+ gert::StorageShape gradUniqueShape = {param.gradUniqueShape, param.gradUniqueShape};
452+ gert::StorageShape uniqueLocalEntryShape = {param.uniqueLocalEntryShape, param.uniqueLocalEntryShape};
453+ gert::StorageShape numUniqueShape = {param.numUniqueShape, param.numUniqueShape};
454+ 
455+ // 7 inputs: commContext, gradFetched, perm, sendCounts, recvCounts, recvLocalEntry, numRecv
456+ std::vector<gert::TilingContextPara::TensorDescription> inputTensorDesc_(
457+ {{commContextShape, param.commContextDtype, param.commContextFormat},
458+ {gradFetchedShape, param.gradFetchedDtype, param.gradFetchedFormat},
459+ {permShape, param.permDtype, ge::FORMAT_ND},
460+ {sendCountsShape, ge::DT_INT32, ge::FORMAT_ND},
461+ {recvCountsShape, ge::DT_INT32, ge::FORMAT_ND},
462+ {recvLocalEntryShape, ge::DT_INT32, ge::FORMAT_ND},
463+ {numUniqueShape, ge::DT_INT32, ge::FORMAT_ND}});
464+ 
465+ // 3 outputs: gradUniqueOut (same dtype as gradFetched), uniqueLocalEntryOut (int32), numUniqueOut (int32)
466+ std::vector<gert::TilingContextPara::TensorDescription> outputTensorDesc_(
467+ {{gradUniqueShape, param.gradUniqueDtype, ge::FORMAT_ND},
468+ {uniqueLocalEntryShape, ge::DT_INT32, ge::FORMAT_ND},
469+ {numUniqueShape, ge::DT_INT32, ge::FORMAT_ND}});
470+ 
471+ std::vector<gert::TilingContextPara::OpAttr> attrs_(
472+ {{"num_entries_per_rank", Ops::Transformer::AnyValue::CreateFrom<int64_t>(param.numEntriesPerRank)},
473+ {"comm_buffer_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(param.commBufferSize)}});
474+ 
475+ return gert::TilingContextPara(OP_NAME, inputTensorDesc_, outputTensorDesc_, attrs_, &compileInfo,
476+ param.socVersion);
atomgit-bot
atomgit-botatomgit-bot7月21日

🟠 High Priority

tiling 测试的参数结构和测试用例存在与 op 定义不匹配的问题:

  1. 缺少第 7 个输入 tensor numRecv(第214-220行):inputTensorDesc_ 只包含 6 个 tensor(commContext, gradFetched, perm, sendCounts, recvCounts, recvLocalEntry),但 op_def(engram_fetch_grad_def.cpp)定义了 7 个输入,包含 numRecv。
  2. 测试用例中缺少 comm_buffer_size attr,却添加了不存在的 hidden_size attr(第228-230行):op_def 中定义的 attr 为 num_entries_per_rank 和 comm_buffer_size,无 hidden_size。测试只提供了 hidden_size 和 num_entries_per_rank。
  3. 测试参数结构 EngramFetchGradTestParam 缺少 numRecvShape、commBufferSize 字段,且所有测试用例 g_testCases 也缺少对应的初始化值。

这将导致 tiling 测试在 TilingCaseExecutor 框架中因输入数量/属性不匹配而失败。

likedislike
不准确?
luozhonglin
luozhonglin
7月24日 评论:
477+}
478+ 
479+TEST_P(EngramFetchGradArch35TilingTest, GeneralCases)
480+{
481+ auto param = GetParam();
482+ auto tilingContextPara = BuildTilingContextPara(param);
483+ ExecuteTestCase(tilingContextPara, param.status, param.expectTilingKey, param.expectTilingData,
484+ param.expectWorkspaces);
485+}
486+ 
487+INSTANTIATE_TEST_CASE_P(EngramFetchGradTilingUT, EngramFetchGradArch35TilingTest, testing::ValuesIn(g_testCases));
488+ 
489+} // namespace EngramFetchGradUT