已合并
all_gather_add通算融合示例算子代码优化 #454
lyt_claire创建于 2025年12月8日
all_gather_add通算融合示例算子代码优化 #454
已合并
共 8 个文件变更+77-78
| @@ -14,8 +14,5 @@ foreach(SUBDIR ${SUBDIRECTORIES}) | |||
| 14 | # 检查子目录中是否存在 CMakeLists.txt | 14 | # 检查子目录中是否存在 CMakeLists.txt |
| 15 | if(EXISTS ${CMAKE_CURRENT_SOURCE_DIR}/${SUBDIR}/CMakeLists.txt) | 15 | if(EXISTS ${CMAKE_CURRENT_SOURCE_DIR}/${SUBDIR}/CMakeLists.txt) |
| 16 | add_subdirectory(${SUBDIR}) | 16 | add_subdirectory(${SUBDIR}) |
| 17 | - if(DEFINED ${SUBDIR}_depends) | ||
| 18 | - set(${SUBDIR}_depends "${${SUBDIR}_depends}" PARENT_SCOPE) | ||
| 19 | - endif() | ||
| 20 | endif() | 17 | endif() |
| 21 | endforeach() | 18 | endforeach() |
| @@ -15,8 +15,4 @@ | |||
| 15 | if(EXISTS ${CMAKE_CURRENT_SOURCE_DIR}/${SUBDIR}/CMakeLists.txt) | 15 | if(EXISTS ${CMAKE_CURRENT_SOURCE_DIR}/${SUBDIR}/CMakeLists.txt) |
| 16 | add_subdirectory(${SUBDIR}) | 16 | add_subdirectory(${SUBDIR}) |
| 17 | endif() | 17 | endif() |
| 18 | - endforeach() | 18 | + endforeach() |
| 19 | - | ||
| 20 | - if (BUILD_OPEN_PROJECT) | ||
| 21 | - get_filename_component(CURRENT_DIR_NAME ${CMAKE_CURRENT_SOURCE_DIR} NAME) | ||
| 22 | -endif() | ||
| @@ -110,6 +110,8 @@ | |||
| 110 | ## 约束说明 | 110 | ## 约束说明 |
| 111 | * 当前该示例算子仅支持固定shape:a(240, 256),b(240 * 2, 256),和固定rank_size = 2。 | 111 | * 当前该示例算子仅支持固定shape:a(240, 256),b(240 * 2, 256),和固定rank_size = 2。 |
| 112 | * 所有输入不支持空tensor场景,取值范围在[-5,5]之间。 | 112 | * 所有输入不支持空tensor场景,取值范围在[-5,5]之间。 |
| 113 | +* 确定性计算: | ||
| 114 | + - allGatherAdd算子默认确定性实现。 | ||
| 113 | ## 调用说明 | 115 | ## 调用说明 |
| 114 | 116 | ||
| 115 | 调用本算子前,请确保已本地下载代码仓,并安装好如下基础依赖、NPU驱动和固件已安装。 | 117 | 调用本算子前,请确保已本地下载代码仓,并安装好如下基础依赖、NPU驱动和固件已安装。 |
| @@ -21,38 +21,26 @@ class AllGatherAdd : public OpDef { | |||
| 21 | this->Input("a") | 21 | this->Input("a") |
| 22 | .ParamType(REQUIRED) | 22 | .ParamType(REQUIRED) |
| 23 | .DataType({ge::DT_FLOAT16}) | 23 | .DataType({ge::DT_FLOAT16}) |
| 24 | - .Format({ge::FORMAT_ND}) | 24 | + .Format({ge::FORMAT_ND}); |
| 25 | - .UnknownShapeFormat({ge::FORMAT_ND}); | ||
| 26 | this->Input("b") | 25 | this->Input("b") |
| 27 | .ParamType(REQUIRED) | 26 | .ParamType(REQUIRED) |
| 28 | .DataType({ge::DT_FLOAT16}) | 27 | .DataType({ge::DT_FLOAT16}) |
| 29 | - .Format({ge::FORMAT_ND}) | 28 | + .Format({ge::FORMAT_ND}); |
| 30 | - .UnknownShapeFormat({ge::FORMAT_ND}); | ||
| 31 | 29 | ||
| 32 | this->Output("c") | 30 | this->Output("c") |
| 33 | .ParamType(REQUIRED) | 31 | .ParamType(REQUIRED) |
| 34 | .DataType({ge::DT_FLOAT16}) | 32 | .DataType({ge::DT_FLOAT16}) |
| 35 | - .Format({ge::FORMAT_ND}) | 33 | + .Format({ge::FORMAT_ND}); |
| 36 | - .UnknownShapeFormat({ge::FORMAT_ND}); | ||
| 37 | this->Output("gather_out") | 34 | this->Output("gather_out") |
| 38 | .ParamType(REQUIRED) | 35 | .ParamType(REQUIRED) |
| 39 | .DataType({ge::DT_FLOAT16}) | 36 | .DataType({ge::DT_FLOAT16}) |
| 40 | - .Format({ge::FORMAT_ND}) | 37 | + .Format({ge::FORMAT_ND}); |
| 41 | - .UnknownShapeFormat({ge::FORMAT_ND}); | ||
| 42 | 38 | ||
| 43 | this->Attr("group").AttrType(REQUIRED).String(); // 通算融合算子属性,表示通信域名称 | 39 | this->Attr("group").AttrType(REQUIRED).String(); // 通算融合算子属性,表示通信域名称 |
| 44 | this->Attr("rank_size").AttrType(REQUIRED).Int(0); | 40 | this->Attr("rank_size").AttrType(REQUIRED).Int(0); |
| 45 | 41 | ||
| 46 | - OpAICoreConfig aicoreConfig; | 42 | + this->AICore().AddConfig("ascend910b"); |
| 47 | - aicoreConfig.DynamicCompileStaticFlag(true) | 43 | + this->AICore().AddConfig("ascend910_93"); |
| 48 | - .DynamicFormatFlag(false) | ||
| 49 | - .DynamicRankSupportFlag(true) | ||
| 50 | - .DynamicShapeSupportFlag(true) | ||
| 51 | - .NeedCheckSupportFlag(false) | ||
| 52 | - .PrecisionReduceFlag(true) | ||
| 53 | - .ExtendCfgInfo("opFile.value", "all_gather_add"); // 这里制定的值会对应到kernel入口文件名.cpp | ||
| 54 | - this->AICore().AddConfig("ascend910b", aicoreConfig); | ||
| 55 | - this->AICore().AddConfig("ascend910_93", aicoreConfig); | ||
| 56 | this->MC2().HcclGroup("group"); // group 属性配置为该算子的通信域名称 | 44 | this->MC2().HcclGroup("group"); // group 属性配置为该算子的通信域名称 |
| 57 | } | 45 | } |
| 58 | }; | 46 | }; |
| @@ -26,7 +26,8 @@ using namespace ge; | |||
| 26 | 26 | ||
| 27 | namespace { | 27 | namespace { |
| 28 | constexpr uint32_t TILE_NUM = 1; | 28 | constexpr uint32_t TILE_NUM = 1; |
| 29 | - constexpr uint32_t COMM_TURN = 1; | 29 | + constexpr uint32_t COMM_TURN = 2; |
| 30 | + const uint32_t WS_SYS_SIZE = 16U * 1024U * 1024U; | ||
| 30 | } | 31 | } |
| 31 | namespace optiling { | 32 | namespace optiling { |
| 32 | 33 | ||
| @@ -59,49 +60,44 @@ static void InitHcclParam(AllGatherAddTilingData* tilingData, const char* group) | |||
| 59 | mc2CcTilingConfig.GetTiling(tilingData->mc2CcTiling); | 60 | mc2CcTilingConfig.GetTiling(tilingData->mc2CcTiling); |
| 60 | } | 61 | } |
| 61 | 62 | ||
| 63 | +ge::graphStatus GetWorkspaceSize(gert::TilingContext* context) | ||
| 64 | +{ | ||
| 65 | + size_t* currentWorkspace = context->GetWorkspaceSizes(1); | ||
| 66 | + OP_CHECK_NULL_WITH_CONTEXT(context, currentWorkspace); | ||
| 67 | + currentWorkspace[0] = WS_SYS_SIZE; | ||
| 68 | + return ge::GRAPH_SUCCESS; | ||
| 69 | +} | ||
| 70 | + | ||
| 62 | static ge::graphStatus AllGatherAddTilingFunc(gert::TilingContext *context) { | 71 | static ge::graphStatus AllGatherAddTilingFunc(gert::TilingContext *context) { |
| 63 | - // 对参数进行校验 | 72 | + // 1.对参数进行校验 |
| 64 | OP_CHECK_IF(AllGatherParamsCheck(context) != ge::GRAPH_SUCCESS, | 73 | OP_CHECK_IF(AllGatherParamsCheck(context) != ge::GRAPH_SUCCESS, |
| 65 | OP_LOGE(context->GetNodeName(), "param is invalid"), return ge::GRAPH_FAILED); | 74 | OP_LOGE(context->GetNodeName(), "param is invalid"), return ge::GRAPH_FAILED); |
| 66 | 75 | ||
| 67 | auto ascendcPlatform = platform_ascendc::PlatformAscendC(context->GetPlatformInfo()); | 76 | auto ascendcPlatform = platform_ascendc::PlatformAscendC(context->GetPlatformInfo()); |
| 68 | context->SetBlockDim(ascendcPlatform.GetCoreNumAiv()); | 77 | context->SetBlockDim(ascendcPlatform.GetCoreNumAiv()); |
| 69 | 78 | ||
| 70 | - // 设置TilingData | 79 | + // 2.获取WorkspaceSize信息 |
| 80 | + OP_CHECK_IF(GetWorkspaceSize(context) != ge::GRAPH_SUCCESS, | ||
| 81 | + OP_LOGE(context, "GetWorkspaceSize error"), return ge::GRAPH_FAILED); | ||
| 82 | + | ||
| 83 | + // 3.设置TilingData | ||
| 71 | AllGatherAddTilingData* tilingData = context->GetTilingData<AllGatherAddTilingData>(); | 84 | AllGatherAddTilingData* tilingData = context->GetTilingData<AllGatherAddTilingData>(); |
| 72 | OP_CHECK_NULL_WITH_CONTEXT(context, tilingData); | 85 | OP_CHECK_NULL_WITH_CONTEXT(context, tilingData); |
| 73 | 86 | ||
| 74 | - tilingData->commTurn = COMM_TURN; | 87 | + tilingData->commTurn = COMM_TURN; // 通信轮次为1时通算串行,大于1时开启通算掩盖 |
| 75 | tilingData->tileNum = TILE_NUM; | 88 | tilingData->tileNum = TILE_NUM; |
| 76 | tilingData->totalElemNum = context->GetInputTensor(1)->GetShapeSize(); | 89 | tilingData->totalElemNum = context->GetInputTensor(1)->GetShapeSize(); |
| 77 | - tilingData->blockElemNum = tilingData->totalElemNum / context->GetBlockDim(); | 90 | + tilingData->blockElemNum = tilingData->totalElemNum / tilingData->commTurn / context->GetBlockDim(); // 每次Add计算只处理前一次通信结果长度的数据 |
| 78 | tilingData->addTileElemNum = tilingData->blockElemNum / tilingData->tileNum; | 91 | tilingData->addTileElemNum = tilingData->blockElemNum / tilingData->tileNum; |
| 79 | uint32_t rankSize = *context->GetAttrs()->GetAttrPointer<uint32_t>(static_cast<int>(1)); | 92 | uint32_t rankSize = *context->GetAttrs()->GetAttrPointer<uint32_t>(static_cast<int>(1)); |
| 80 | - tilingData->gatherTileElemNum = tilingData->totalElemNum / rankSize; | 93 | + tilingData->addCoresPerRank = context->GetBlockDim() / rankSize; // 进行Add计算之前需要根据每个rank分到的核数来判断当前核的计算地址偏移 |
| 81 | - | 94 | + tilingData->gatherTileElemNum = tilingData->totalElemNum / rankSize / tilingData->commTurn; // 每轮通信的数据长度 |
L | |||
| 82 | - // 设置workspaceSize gather out需要额外的临时内存,大小与b输入一致 | ||
| 83 | - size_t* currentWorkspace = context->GetWorkspaceSizes(1); | ||
| 84 | - OP_CHECK_NULL_WITH_CONTEXT(context,currentWorkspace); | ||
| 85 | - // 如需使用系统workspace需要调用GetLibApiWorkSpaceSize获取系统workspace大小 | ||
| 86 | - uint32_t sysWorkSpaceSize = ascendcPlatform.GetLibApiWorkSpaceSize(); | ||
| 87 | - // 预留18M + gather_out | ||
| 88 | - auto dataType = context->GetInputTensor(0)->GetDataType(); | ||
| 89 | - currentWorkspace[0] = sysWorkSpaceSize + tilingData->totalElemNum * sizeof(dataType); | ||
| 90 | 95 | ||
| 91 | auto group = context->GetAttrs()->GetAttrPointer<char>(static_cast<int>(0)); | 96 | auto group = context->GetAttrs()->GetAttrPointer<char>(static_cast<int>(0)); |
| 92 | InitHcclParam(tilingData, group); | 97 | InitHcclParam(tilingData, group); |
| 93 | return ge::GRAPH_SUCCESS; | 98 | return ge::GRAPH_SUCCESS; |
| 94 | } | 99 | } |
| 95 | 100 | ||
| 96 | -struct AllGatherAddCompileInfo {}; | ||
| 97 | - | ||
| 98 | -static ge::graphStatus TilingParseForAllGatherAdd([[maybe_unused]] gert::TilingParseContext *context) | ||
| 99 | -{ | ||
| 100 | - (void)context; | ||
| 101 | - return ge::GRAPH_SUCCESS; | ||
| 102 | -} | ||
| 103 | - | ||
| 104 | IMPL_OP_OPTILING(AllGatherAdd) | 101 | IMPL_OP_OPTILING(AllGatherAdd) |
| 105 | - .Tiling(AllGatherAddTilingFunc) | 102 | + .Tiling(AllGatherAddTilingFunc); |
| 106 | - .TilingParse<AllGatherAddCompileInfo>(TilingParseForAllGatherAdd); | ||
| 107 | } // namespace optiling | 103 | } // namespace optiling |
| @@ -20,16 +20,14 @@ using namespace AscendC; | |||
| 20 | extern "C" __global__ __aicore__ void all_gather_add(GM_ADDR aGM, GM_ADDR bGM, GM_ADDR cGM, | 20 | extern "C" __global__ __aicore__ void all_gather_add(GM_ADDR aGM, GM_ADDR bGM, GM_ADDR cGM, |
| 21 | GM_ADDR gatherGM, GM_ADDR workspaceGM, GM_ADDR tilingGM) | 21 | GM_ADDR gatherGM, GM_ADDR workspaceGM, GM_ADDR tilingGM) |
| 22 | { | 22 | { |
| 23 | + // 设置kernel类型,AIC、AIV混合场景下,控制算子执行时仅启动AI Core上的Vector核 | ||
| 23 | KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_MIX_AIV_1_0); | 24 | KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_MIX_AIV_1_0); |
| 24 | // 注册算子Tiling结构体 | 25 | // 注册算子Tiling结构体 |
| 25 | REGISTER_TILING_DEFAULT(AllGatherAddTilingData); | 26 | REGISTER_TILING_DEFAULT(AllGatherAddTilingData); |
| 26 | - auto tiling = (__gm__ AllGatherAddTilingData*)tilingGM; | ||
| 27 | GET_TILING_DATA(tilingData, tilingGM); | 27 | GET_TILING_DATA(tilingData, tilingGM); |
| 28 | - | ||
| 29 | TPipe pipe; | 28 | TPipe pipe; |
| 30 | - GM_ADDR contextGM = GetHcclContext<HCCL_GROUP_ID_0>(); | ||
| 31 | 29 | ||
| 32 | AllGatherAdd allGatherAdd; | 30 | AllGatherAdd allGatherAdd; |
| 33 | - allGatherAdd.Init(aGM, bGM, cGM, gatherGM, workspaceGM, contextGM, &tilingData, &pipe); | 31 | + allGatherAdd.Init(aGM, bGM, cGM, gatherGM, &tilingData, &pipe); |
| 34 | allGatherAdd.Process(); | 32 | allGatherAdd.Process(); |
| 35 | } | 33 | } |
| @@ -19,18 +19,19 @@ | |||
| 19 | 19 | ||
| 20 | 20 | ||
| 21 | 21 | ||
| 22 | -constexpr int32_t ALLGATHER_ADD_BUFFER_NUM = 1; | 22 | +constexpr int32_t ADD_BUFFER_NUM = 2; |
| 23 | 23 | ||
| 24 | namespace AscendC { | 24 | namespace AscendC { |
| 25 | class AllGatherAdd { | 25 | class AllGatherAdd { |
| 26 | public: | 26 | public: |
| 27 | __aicore__ inline AllGatherAdd(){}; | 27 | __aicore__ inline AllGatherAdd(){}; |
| 28 | __aicore__ inline void Init(GM_ADDR aGM, GM_ADDR bGM, GM_ADDR cGM, GM_ADDR gatherGM, | 28 | __aicore__ inline void Init(GM_ADDR aGM, GM_ADDR bGM, GM_ADDR cGM, GM_ADDR gatherGM, |
| 29 | - GM_ADDR workspaceGM, GM_ADDR contextGM, AllGatherAddTilingData *tilingData, TPipe *tPipe); | 29 | + AllGatherAddTilingData *tilingData, TPipe *tPipe); |
| 30 | __aicore__ inline void Process(); | 30 | __aicore__ inline void Process(); |
| 31 | 31 | ||
| 32 | private: | 32 | private: |
| 33 | __aicore__ inline void HcclPrepare(); | 33 | __aicore__ inline void HcclPrepare(); |
| 34 | + __aicore__ inline void CalcAddGmAddr(int32_t commTurn); | ||
| 34 | __aicore__ inline void CopyIn(int32_t progress); | 35 | __aicore__ inline void CopyIn(int32_t progress); |
| 35 | __aicore__ inline void CopyOut(int32_t progress); | 36 | __aicore__ inline void CopyOut(int32_t progress); |
| 36 | __aicore__ inline void Compute(); | 37 | __aicore__ inline void Compute(); |
| @@ -40,14 +41,17 @@ private: | |||
| 40 | 41 | ||
| 41 | AllGatherAddTilingData *tilingData_; | 42 | AllGatherAddTilingData *tilingData_; |
| 42 | 43 | ||
| 43 | - TPipe *tPipe_; | ||
| 44 | Hccl<HCCL_SERVER_TYPE_AICPU> hccl_; | 44 | Hccl<HCCL_SERVER_TYPE_AICPU> hccl_; |
| 45 | 45 | ||
| 46 | - TQue<QuePosition::VECIN, ALLGATHER_ADD_BUFFER_NUM> inputQueueGather; | 46 | + TQue<QuePosition::VECIN, ADD_BUFFER_NUM> inputQueueGather; |
| 47 | - TQue<QuePosition::VECIN, ALLGATHER_ADD_BUFFER_NUM> inputQueueB; | 47 | + TQue<QuePosition::VECIN, ADD_BUFFER_NUM> inputQueueB; |
| 48 | - TQue<QuePosition::VECOUT, ALLGATHER_ADD_BUFFER_NUM> outputQueueC; | 48 | + TQue<QuePosition::VECOUT, ADD_BUFFER_NUM> outputQueueC; |
| 49 | + | ||
| 50 | + GM_ADDR aGM_; | ||
| 51 | + GM_ADDR bGM_; | ||
| 52 | + GM_ADDR cGM_; | ||
| 53 | + GM_ADDR gatherGM_; | ||
| 49 | 54 | ||
| 50 | - GlobalTensor<half> inputAGM; | ||
| 51 | GlobalTensor<half> gatherOutGM; | 55 | GlobalTensor<half> gatherOutGM; |
| 52 | GlobalTensor<half> inputBGM; | 56 | GlobalTensor<half> inputBGM; |
| 53 | GlobalTensor<half> outputCGM; | 57 | GlobalTensor<half> outputCGM; |
| @@ -55,39 +59,42 @@ private: | |||
| 55 | int64_t blockElemNum_ = 0; | 59 | int64_t blockElemNum_ = 0; |
| 56 | int64_t tileNum_ = 0; | 60 | int64_t tileNum_ = 0; |
| 57 | uint32_t addTileElemNum_ = 0; | 61 | uint32_t addTileElemNum_ = 0; |
| 62 | + int64_t blockIdx_ = 0; | ||
| 63 | + uint64_t elemNumPerRank_ = 0; | ||
| 58 | 64 | ||
| 59 | HcclHandle handleId_{ INVALID_HANDLE_ID }; | 65 | HcclHandle handleId_{ INVALID_HANDLE_ID }; |
| 60 | }; | 66 | }; |
| 61 | 67 | ||
| 62 | __aicore__ inline void AllGatherAdd::Init(GM_ADDR aGM, GM_ADDR bGM, GM_ADDR cGM, GM_ADDR gatherGM, | 68 | __aicore__ inline void AllGatherAdd::Init(GM_ADDR aGM, GM_ADDR bGM, GM_ADDR cGM, GM_ADDR gatherGM, |
| 63 | - GM_ADDR workspaceGM, GM_ADDR contextGM, AllGatherAddTilingData *tilingData, TPipe *tPipe) | 69 | + AllGatherAddTilingData *tilingData, TPipe *tPipe) |
| 64 | { | 70 | { |
| 71 | + aGM_ = aGM; | ||
| 72 | + bGM_ = bGM; | ||
| 73 | + cGM_ = cGM; | ||
| 74 | + gatherGM_ = gatherGM; | ||
| 75 | + | ||
| 65 | tilingData_ = tilingData; | 76 | tilingData_ = tilingData; |
| 66 | - tPipe_ = tPipe; | ||
| 67 | blockElemNum_ = tilingData->blockElemNum; | 77 | blockElemNum_ = tilingData->blockElemNum; |
| 68 | - addTileElemNum_ = tilingData->addTileElemNum; | 78 | + addTileElemNum_ = tilingData->addTileElemNum / ADD_BUFFER_NUM; |
| 69 | tileNum_ = tilingData->tileNum; | 79 | tileNum_ = tilingData->tileNum; |
| 80 | + blockIdx_ = AscendC::GetBlockIdx(); | ||
| 70 | 81 | ||
| 71 | // 初始化hccl对象 | 82 | // 初始化hccl对象 |
| 83 | + GM_ADDR contextGM = GetHcclContext<HCCL_GROUP_ID_0>(); | ||
| 72 | hccl_.InitV2(contextGM, tilingData); | 84 | hccl_.InitV2(contextGM, tilingData); |
| 73 | hccl_.SetCcTilingV2(offsetof(AllGatherAddTilingData, mc2CcTiling)); | 85 | hccl_.SetCcTilingV2(offsetof(AllGatherAddTilingData, mc2CcTiling)); |
| 74 | - | ||
| 75 | - // 传入全局数据的指针,并设置存储大小 | ||
| 76 | - inputAGM.SetGlobalBuffer((__gm__ half*)aGM, tilingData->gatherTileElemNum); // 非多轮切分AllGather场景,每张卡参与Gather的数据大小为{240,256} | ||
| 77 | - gatherOutGM.SetGlobalBuffer((__gm__ half*)gatherGM + blockElemNum_ * AscendC::GetBlockIdx(), blockElemNum_); | ||
| 78 | - inputBGM.SetGlobalBuffer((__gm__ half*)bGM + blockElemNum_ * AscendC::GetBlockIdx(), blockElemNum_); | ||
| 79 | - outputCGM.SetGlobalBuffer((__gm__ half*)cGM + blockElemNum_ * AscendC::GetBlockIdx(), blockElemNum_); | ||
| 80 | 86 | ||
| 81 | - tPipe_->InitBuffer(inputQueueGather, ALLGATHER_ADD_BUFFER_NUM, addTileElemNum_ * sizeof(half)); | 87 | + tPipe->InitBuffer(inputQueueGather, ADD_BUFFER_NUM, addTileElemNum_ * sizeof(half)); |
| 82 | - tPipe_->InitBuffer(inputQueueB, ALLGATHER_ADD_BUFFER_NUM, addTileElemNum_ * sizeof(half)); | 88 | + tPipe->InitBuffer(inputQueueB, ADD_BUFFER_NUM, addTileElemNum_ * sizeof(half)); |
| 83 | - tPipe_->InitBuffer(outputQueueC, ALLGATHER_ADD_BUFFER_NUM, addTileElemNum_ * sizeof(half)); | 89 | + tPipe->InitBuffer(outputQueueC, ADD_BUFFER_NUM, addTileElemNum_ * sizeof(half)); |
| 84 | } | 90 | } |
| 85 | 91 | ||
| 86 | __aicore__ inline void AllGatherAdd::HcclPrepare() | 92 | __aicore__ inline void AllGatherAdd::HcclPrepare() |
| 87 | { | 93 | { |
| 94 | + elemNumPerRank_ = tilingData_->gatherTileElemNum * tilingData_->commTurn; // 通信多轮切分,多张卡的数据拼接到gatherOutGM时,相邻数据块的起始地址偏移 | ||
| 88 | // 下发通信任务 | 95 | // 下发通信任务 |
| 89 | - handleId_ = hccl_.AllGather<true>((__gm__ uint8_t*)this->inputAGM.GetPhyAddr(), (__gm__ uint8_t*)this->gatherOutGM.GetPhyAddr(), tilingData_->gatherTileElemNum, | 96 | + handleId_ = hccl_.AllGather<true>(aGM_, gatherGM_, tilingData_->gatherTileElemNum, |
| 90 | - HcclDataType::HCCL_DATA_TYPE_FP16, 0, tilingData_->commTurn); | 97 | + HcclDataType::HCCL_DATA_TYPE_FP16, elemNumPerRank_, tilingData_->commTurn); |
| 91 | } | 98 | } |
| 92 | 99 | ||
| 93 | __aicore__ inline void AllGatherAdd::CopyIn(int32_t progress) | 100 | __aicore__ inline void AllGatherAdd::CopyIn(int32_t progress) |
| @@ -124,12 +131,26 @@ __aicore__ inline void AllGatherAdd::HcclFinalize() | |||
| 124 | hccl_.Finalize(); | 131 | hccl_.Finalize(); |
| 125 | } | 132 | } |
| 126 | 133 | ||
| 134 | +__aicore__ inline void AllGatherAdd::CalcAddGmAddr(int32_t commTurn) | ||
| 135 | +{ | ||
| 136 | + uint32_t commOffset = commTurn * tilingData_->gatherTileElemNum; // 1.根据通信轮次偏移单个通信数据大小 | ||
| 137 | + uint32_t blockOffset = blockIdx_ / tilingData_->addCoresPerRank * elemNumPerRank_; // 2.根据rank数和aivId判断当前核被分到处理哪个rank的通信数据 | ||
| 138 | + uint32_t totalOffset = commOffset + blockOffset + (blockIdx_ % tilingData_->addCoresPerRank) * blockElemNum_; // 3.最终偏移需要再加上当前核在所处理rank数据上的偏移 | ||
| 139 | + gatherOutGM.SetGlobalBuffer((__gm__ half*)gatherGM_ + totalOffset, blockElemNum_); | ||
| 140 | + inputBGM.SetGlobalBuffer((__gm__ half*)bGM_ + totalOffset, blockElemNum_); | ||
| 141 | + outputCGM.SetGlobalBuffer((__gm__ half*)cGM_ + totalOffset, blockElemNum_); | ||
| 142 | +} | ||
| 143 | + | ||
| 127 | __aicore__ inline void AllGatherAdd::Process() | 144 | __aicore__ inline void AllGatherAdd::Process() |
| 128 | { | 145 | { |
| 129 | HcclPrepare(); | 146 | HcclPrepare(); |
| 147 | + int addLoop = tileNum_ * ADD_BUFFER_NUM; | ||
| 130 | for (int i = 0; i < tilingData_->commTurn; i++) { | 148 | for (int i = 0; i < tilingData_->commTurn; i++) { |
| 131 | - hccl_.Wait(handleId_); | 149 | + hccl_.Wait(handleId_); |
| 132 | - for (int j = 0; j < tileNum_; j++) { | 150 | + // 根据通信轮次和rankSize计算本核需要处理数据的起始地址 |
| 151 | + CalcAddGmAddr(i); | ||
| 152 | + // 对前一轮的通信结果进行Add计算 | ||
| 153 | + for (int j = 0; j < addLoop; j++) { | ||
| 133 | CopyIn(j); | 154 | CopyIn(j); |
| 134 | Compute(); | 155 | Compute(); |
| 135 | CopyOut(j); | 156 | CopyOut(j); |
| @@ -28,6 +28,7 @@ struct AllGatherAddTilingData { | |||
| 28 | uint32_t tileNum; | 28 | uint32_t tileNum; |
| 29 | uint32_t addTileElemNum; | 29 | uint32_t addTileElemNum; |
| 30 | uint32_t gatherTileElemNum; | 30 | uint32_t gatherTileElemNum; |
| 31 | + uint32_t addCoresPerRank; | ||
| 31 | }; | 32 | }; |
| 32 | 33 | ||
| 33 | 34 | ||
workspaceSize需要申请