已合并
all_gather_add通算融合示例算子代码优化 #454
all_gather_add通算融合示例算子代码优化 #454
已合并
lyt_claire创建于 2025年12月8日
8 个文件变更+77-78
@@ -14,8 +14,5 @@ foreach(SUBDIR ${SUBDIRECTORIES})
14 # 检查子目录中是否存在 CMakeLists.txt14 # 检查子目录中是否存在 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()
21endforeach()18endforeach()
@@ -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 
27namespace {27namespace {
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}
31namespace optiling {32namespace 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+ 
62static ge::graphStatus AllGatherAddTilingFunc(gert::TilingContext *context) {71static 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- // 设置TilingData79+ // 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
Llyt_claire1月28日

workspaceSize需要申请

likedislike
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- 
104IMPL_OP_OPTILING(AllGatherAdd)101IMPL_OP_OPTILING(AllGatherAdd)
105- .Tiling(AllGatherAddTilingFunc)102+ .Tiling(AllGatherAddTilingFunc);
106- .TilingParse<AllGatherAddCompileInfo>(TilingParseForAllGatherAdd);
107} // namespace optiling103} // namespace optiling
@@ -20,16 +20,14 @@ using namespace AscendC;
20extern "C" __global__ __aicore__ void all_gather_add(GM_ADDR aGM, GM_ADDR bGM, GM_ADDR cGM,20extern "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#include "kernel_tiling/kernel_tiling.h"19#include "kernel_tiling/kernel_tiling.h"
20#include "all_gather_add_tiling.h"20#include "all_gather_add_tiling.h"
21 21 
22-constexpr int32_t ALLGATHER_ADD_BUFFER_NUM = 1;22+constexpr int32_t ADD_BUFFER_NUM = 2;
23 23 
24namespace AscendC {24namespace AscendC {
25class AllGatherAdd {25class AllGatherAdd {
26public:26public:
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 
32private:32private:
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#endif //ALL_GATHER_ADD_TILING_H34#endif //ALL_GATHER_ADD_TILING_H