| @@ -24,7 +24,9 @@ | |||
| 24 | 24 | ||
| 25 | using namespace op; | 25 | using 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 | |||
| 42 | extern "C" { | 49 | extern "C" { |
| 43 | 50 | ||
| 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 { | |||
| 20 | public: | 20 | public: |
| 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 @@ | |||
| 26 | using namespace AscendC; | 26 | using namespace AscendC; |
| 27 | using namespace ge; | 27 | using namespace ge; |
| 28 | 28 | ||
| 29 | -namespace MC2Tiling { | 29 | +namespace Mc2Tiling { |
| 30 | constexpr uint32_t COMM_CONTEXT_INDEX = 0U; | 30 | constexpr uint32_t COMM_CONTEXT_INDEX = 0U; |
| 31 | constexpr uint32_t INDICES_INDEX = 1U; | 31 | constexpr uint32_t INDICES_INDEX = 1U; |
| 32 | +constexpr uint32_t LOCAL_STORAGE_ADDR_INDEX = 2U; | ||
| 32 | constexpr uint32_t FETCHED_INDEX = 0U; | 33 | constexpr 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 | ||
| 34 | constexpr uint32_t ATTR_HIDDEN_SIZE_INDEX = 0U; | 40 | constexpr uint32_t ATTR_HIDDEN_SIZE_INDEX = 0U; |
| 35 | constexpr uint32_t ATTR_NUM_ENTRIES_PER_RANK_INDEX = 1U; | 41 | constexpr 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 | ||
| 37 | constexpr uint32_t DIM_ONE = 1U; | 46 | constexpr uint32_t DIM_ONE = 1U; |
| 38 | constexpr uint32_t DIM_TWO = 2U; | 47 | constexpr uint32_t DIM_TWO = 2U; |
| 39 | constexpr uint32_t SYSTEM_NEED_WORKSPACE = 16U * 1024 * 1024; | 48 | constexpr uint32_t SYSTEM_NEED_WORKSPACE = 16U * 1024 * 1024; |
| 40 | 49 | ||
| 41 | constexpr int32_t HIDDEN_SIZE_ALIGN = 128; | 50 | constexpr 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 | ||
| 43 | static const std::vector<ge::DataType> OUTPUT_DTYPE_LIST = {ge::DT_BF16, ge::DT_FLOAT16, ge::DT_FLOAT}; | 65 | static 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 | |||
| 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 key | 550 | * @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 info | 635 | // 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 data | 639 | // 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 key | 643 | // 5. set tiling key |
| 418 | - SetTilingKey(context); | 644 | + SetTilingKey(context, isTraining); |
| 419 | 645 | ||
| 420 | // 6. set workspace | 646 | // 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 info | 650 | // 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 | ||
| 430 | IMPL_OP_OPTILING(EngramFetch).Tiling(EngramFetchTilingFunc); | 656 | IMPL_OP_OPTILING(EngramFetch).Tiling(EngramFetchTilingFunc); |
| 431 | -} // namespace MC2Tiling | 657 | +} // namespace Mc2Tiling |
| @@ -26,7 +26,9 @@ | |||
| 26 | using namespace Mc2Kernel; | 26 | using namespace Mc2Kernel; |
| 27 | 27 | ||
| 28 | template <uint32_t EngramFetchMode> | 28 | template <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 | 33 | ||
| @@ -19,14 +19,14 @@ | |||
| 19 | 19 | ||
| 20 | 20 | ||
| 21 | 21 | ||
| 22 | + | ||
| 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是否合法 |
| 29 | ASCENDC_TPL_SEL(ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(EngramFetchMode, ASCENDC_TPL_UI_LIST, | 29 | ASCENDC_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 | 32 | ||
| @@ -11,6 +11,18 @@ | |||
| 11 | /*! | 11 | /*! |
| 12 | * \file test_aclnn_engram_fetch.cpp | 12 | * \file test_aclnn_engram_fetch.cpp |
| 13 | * \brief engram_fetch 算子 op_api 侧 aclnn 接口 UT | 13 | * \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 | 28 | ||
| @@ -21,8 +33,12 @@ | |||
| 21 | 33 | ||
| 22 | 34 | ||
| 23 | extern "C" { | 35 | extern "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 | ||
| 28 | aclnnStatus aclnnEngramFetch(void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, aclrtStream stream); | 44 | aclnnStatus 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 | ||
| 67 | TEST_F(AclnnEngramFetchTest, ascend950_nullptr_commContext) | 118 | TEST_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 | + | ||
| 17 | + | ||
| 18 | + | ||
| 19 | + | ||
| 20 | + | ||
| 21 | + | ||
| 22 | + | ||
| 23 | + | ||
| 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 | + | ||
| 57 | +extern "C" { | ||
| 58 | + | ||
| 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 | + | ||
| 84 | +} | ||
| 85 | + | ||
| @@ -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 | + | ||
| 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 | + | ||
| 23 | + | ||
| 24 | + | ||
| 25 | + | ||
| 26 | + | ||
| 27 | + | ||
| 28 | + | ||
| 29 | + | ||
| 30 | + | ||
| 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 | + | ||
| 27 | + | ||
| 28 | + | ||
| 29 | + | ||
| 30 | + | ||
| 31 | + | ||
| 32 | + | ||
| 33 | + | ||
| 34 | + | ||
| 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 | +} | ||
🟡 Medium Priority 两个 kernel 入口函数 后果:
建议:完成 kernel 实现后再提交,或如果是有意为之的 WIP 提交,将 kernel 函数体替换为明确的错误返回(如 ![]() ![]() 不准确? | |||
| @@ -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 | + | ||
| 17 | + | ||
| 18 | + | ||
| 19 | + | ||
| 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 | + | ||
| @@ -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 | + | ||
| 17 | + | ||
| 18 | + | ||
| 19 | + | ||
| 20 | + | ||
| 21 | + | ||
| 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 | + | ||
| @@ -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 | + | ||
| 26 | + | ||
| 27 | + | ||
| 28 | + | ||
| 29 | + | ||
| 30 | + | ||
| 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 | +} | ||
🔴 Critical 测试文件 这导致测试中的 建议:将 extern 声明改为与实际实现签名一致:在 recvLocalEntry 之后增加 ![]() ![]() 不准确? | |||
| 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 | + | ||
| 26 | + | ||
| 27 | + | ||
| 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 ¶m) | ||
| 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 ¶m) | ||
| 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); | ||
🟠 High Priority tiling 测试的参数结构和测试用例存在与 op 定义不匹配的问题:
这将导致 tiling 测试在 TilingCaseExecutor 框架中因输入数量/属性不匹配而失败。 ![]() ![]() 不准确? | |||
| 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 | ||


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