已开启
AicpuAllGatherConcurMeshNHR 算法适配 #5455
AicpuAllGatherConcurMeshNHR 算法适配 #5455
已开启
hulida创建于 1 天前
3 个文件变更+23-13
@@ -23,7 +23,11 @@
23 23 
24namespace mc2_ops_hccl {24namespace mc2_ops_hccl {
25 25 
26-constexpr u32 CLOS_PORT_NUM = 4;26+constexpr u32 CLOS_BW = 10;
27+constexpr u32 MESH_BW = 11;
28+constexpr u32 CLOS_JETTY = 4;
29+constexpr u32 MESH_BW_AICPU = 37;
30+constexpr u32 CLOS_BW_AICPU = 25;
27 31 
28template <typename AlgTopoMatch, typename InsAlgTemplate0, typename InsAlgTemplate1>32template <typename AlgTopoMatch, typename InsAlgTemplate0, typename InsAlgTemplate1>
29InsV2AllGatherConcurrentExecutor<AlgTopoMatch, InsAlgTemplate0, InsAlgTemplate1>::InsV2AllGatherConcurrentExecutor()33InsV2AllGatherConcurrentExecutor<AlgTopoMatch, InsAlgTemplate0, InsAlgTemplate1>::InsV2AllGatherConcurrentExecutor()
@@ -105,10 +109,7 @@ HcclResult InsV2AllGatherConcurrentExecutor<AlgTopoMatch, InsAlgTemplate0, InsAl
105 CommTopo temp0PriorityTopo = COMM_TOPO_1DMESH;109 CommTopo temp0PriorityTopo = COMM_TOPO_1DMESH;
106 CHK_RET(CalcChannelRequestMesh1DWithPriorityTopo(110 CHK_RET(CalcChannelRequestMesh1DWithPriorityTopo(
107 comm, param, topoInfo, temp0HierarchyInfo, temp0Channels, temp0PriorityTopo));111 comm, param, topoInfo, temp0HierarchyInfo, temp0Channels, temp0PriorityTopo));
108- CommTopo temp1PriorityTopo = COMM_TOPO_CLOS;112+ CHK_RET(CalcChannelRequestNhrMultiJetty(comm, param, topoInfo, temp1HierarchyInfo, temp1Channels));
109- CHK_RET(CalcChannelRequestNHRWithPriorityTopo(
110- comm, param, topoInfo, temp1HierarchyInfo, temp1Channels, temp1PriorityTopo));
111- 
112 CHK_PRT_RET(113 CHK_PRT_RET(
113 temp0Channels.size() != temp1Channels.size(),114 temp0Channels.size() != temp1Channels.size(),
114 HCCL_ERROR(115 HCCL_ERROR(
@@ -116,6 +117,7 @@ HcclResult InsV2AllGatherConcurrentExecutor<AlgTopoMatch, InsAlgTemplate0, InsAl
116 "temp1Channels.size()[%zu]",117 "temp1Channels.size()[%zu]",
117 temp0Channels.size(), temp1Channels.size()),118 temp0Channels.size(), temp1Channels.size()),
118 HcclResult::HCCL_E_INTERNAL);119 HcclResult::HCCL_E_INTERNAL);
120+ resourceRequest.channels.resize(1);
119 resourceRequest.channels[0].insert(121 resourceRequest.channels[0].insert(
120 resourceRequest.channels[0].end(), temp0Channels.begin(), temp0Channels.end());122 resourceRequest.channels[0].end(), temp0Channels.begin(), temp0Channels.end());
121 resourceRequest.channels[0].insert(123 resourceRequest.channels[0].insert(
@@ -138,6 +140,7 @@ void InsV2AllGatherConcurrentExecutor<AlgTopoMatch, InsAlgTemplate0, InsAlgTempl
138 tempAlgParams.buffInfo.outBuffType = BufferType::OUTPUT;140 tempAlgParams.buffInfo.outBuffType = BufferType::OUTPUT;
139 tempAlgParams.count = dataCountPerLoop;141 tempAlgParams.count = dataCountPerLoop;
140 tempAlgParams.sliceSize = dataCountPerLoop * dataTypeSize_;142 tempAlgParams.sliceSize = dataCountPerLoop * dataTypeSize_;
143+ tempAlgParams.tailSize = dataCountPerLoop * dataTypeSize_;
141 tempAlgParams.buffInfo.inBuffBaseOff = dataOffset;144 tempAlgParams.buffInfo.inBuffBaseOff = dataOffset;
142 tempAlgParams.buffInfo.outBuffBaseOff = dataOffset;145 tempAlgParams.buffInfo.outBuffBaseOff = dataOffset;
143 tempAlgParams.buffInfo.hcclBuffBaseOff = scratchOffset;146 tempAlgParams.buffInfo.hcclBuffBaseOff = scratchOffset;
@@ -159,10 +162,17 @@ void InsV2AllGatherConcurrentExecutor<AlgTopoMatch, InsAlgTemplate0, InsAlgTempl
159 162 
160template <typename AlgTopoMatch, typename InsAlgTemplate0, typename InsAlgTemplate1>163template <typename AlgTopoMatch, typename InsAlgTemplate0, typename InsAlgTemplate1>
161void InsV2AllGatherConcurrentExecutor<AlgTopoMatch, InsAlgTemplate0, InsAlgTemplate1>::GetParallelDataSplit(164void InsV2AllGatherConcurrentExecutor<AlgTopoMatch, InsAlgTemplate0, InsAlgTemplate1>::GetParallelDataSplit(
162- std::vector<float>& splitDataSize) const165+ const OpParam& param, std::vector<float>& splitDataSize) const
163{166{
164- const u32 portNum0 = rankSize_ - 1; // mesh端口数为rank size - 1167+ u32 portNum0 = rankSize_ - 1; // mesh端口数为rank size - 1
165- const u32 portNum1 = CLOS_PORT_NUM;168+ u32 portNum1 = CLOS_JETTY;
169+ if (param.engine == CommEngine::COMM_ENGINE_CCU) {
170+ portNum0 = MESH_BW;
171+ portNum1 = CLOS_BW;
172+ } else if (param.opExecuteConfig == OpExecuteConfig::AICPU_TS) {
173+ portNum0 = MESH_BW_AICPU;
174+ portNum1 = CLOS_BW_AICPU;
175+ }
166 double splitData = static_cast<double>(portNum0) / (portNum0 + portNum1);176 double splitData = static_cast<double>(portNum0) / (portNum0 + portNum1);
167 splitDataSize.push_back(splitData);177 splitDataSize.push_back(splitData);
168 splitDataSize.push_back(1 - splitData);178 splitDataSize.push_back(1 - splitData);
@@ -287,7 +297,7 @@ HcclResult InsV2AllGatherConcurrentExecutor<AlgTopoMatch, InsAlgTemplate0, InsAl
287 297 
288 // 计算数据切分比例298 // 计算数据切分比例
289 std::vector<float> dataSplitSize;299 std::vector<float> dataSplitSize;
290- GetParallelDataSplit(dataSplitSize);300+ GetParallelDataSplit(param, dataSplitSize);
291 301 
292 // 缓存切分302 // 缓存切分
293 u32 scratchMultiplierforTemp0 =303 u32 scratchMultiplierforTemp0 =
@@ -303,7 +313,7 @@ HcclResult InsV2AllGatherConcurrentExecutor<AlgTopoMatch, InsAlgTemplate0, InsAl
303 scratchMemBlockSize = (maxTmpMemSize_ / HCCL_MIN_SLICE_ALIGN / totalScratchMultiple) * HCCL_MIN_SLICE_ALIGN;313 scratchMemBlockSize = (maxTmpMemSize_ / HCCL_MIN_SLICE_ALIGN / totalScratchMultiple) * HCCL_MIN_SLICE_ALIGN;
304 }314 }
305 u64 scratchSizeforTemp0 = ScratchMultiplier0 * scratchMemBlockSize;315 u64 scratchSizeforTemp0 = ScratchMultiplier0 * scratchMemBlockSize;
306- u64 scratchSizeforTemp1 = scratchMemBlockSize - scratchSizeforTemp0;316+ u64 scratchSizeforTemp1 = maxTmpMemSize_ - scratchSizeforTemp0;
307 u64 scratchOffsetforTemp0 = 0;317 u64 scratchOffsetforTemp0 = 0;
308 u64 scratchOffsetforTemp1 = scratchSizeforTemp0;318 u64 scratchOffsetforTemp1 = scratchSizeforTemp0;
309 319 
@@ -389,7 +399,7 @@ HcclResult InsV2AllGatherConcurrentExecutor<AlgTopoMatch, InsAlgTemplate0, InsAl
389 399 
390// 算法注册400// 算法注册
391REGISTER_EXECUTOR_BY_TWO_TEMPS(401REGISTER_EXECUTOR_BY_TWO_TEMPS(
392- HcclCMDType::HCCL_CMD_ALLGATHER, InsAllGatherConcurrentMesh1DNHR, InsV2AllGatherConcurrentExecutor, TopoMatchUBX,402+ HcclCMDType::HCCL_CMD_ALLGATHER, AicpuAllGatherConcurMeshNHR, InsV2AllGatherConcurrentExecutor, TopoMatchUBX,
393 InsTempAllGatherMesh1D, InsTempAllGatherNHR);403 InsTempAllGatherMesh1D, InsTempAllGatherNHR);
394 404 
395#if !defined(AICPU_COMPILE) && MC2_CLIENT_ENABLE_CCU405#if !defined(AICPU_COMPILE) && MC2_CLIENT_ENABLE_CCU
@@ -41,7 +41,7 @@ private:
41 const OpParam& param, const TopoInfoWithNetLayerDetails* topoInfo,41 const OpParam& param, const TopoInfoWithNetLayerDetails* topoInfo,
42 const AlgHierarchyInfoForAllLevel& algHierarchyInfo);42 const AlgHierarchyInfoForAllLevel& algHierarchyInfo);
43 43 
44- void GetParallelDataSplit(std::vector<float>& splitDataSize) const;44+ void GetParallelDataSplit(const OpParam& param, std::vector<float>& splitDataSize) const;
45 45 
46 void GenTemplateAlgParams(46 void GenTemplateAlgParams(
47 const OpParam& param, const AlgResourceCtxSerializable& resCtx, const u64 dataOffset,47 const OpParam& param, const AlgResourceCtxSerializable& resCtx, const u64 dataOffset,
@@ -175,7 +175,7 @@ SelectorStatus AllGatherAutoSelector::SelectAicpuAlgo(
175 HCCL_ERROR("[AllGatherAutoSelector] CheckClosNumMultipleOfMeshNum failed."), SelectorStatus::NOT_MATCH);175 HCCL_ERROR("[AllGatherAutoSelector] CheckClosNumMultipleOfMeshNum failed."), SelectorStatus::NOT_MATCH);
176 if (isMeshNumEqualToClosNum && (topoInfo->userRankSize <= MAX_RANK_NUM_FOR_CONCURRENT_ALGO)) {176 if (isMeshNumEqualToClosNum && (topoInfo->userRankSize <= MAX_RANK_NUM_FOR_CONCURRENT_ALGO)) {
177 if (dataSize > SMALL_COUNT_512KB) {177 if (dataSize > SMALL_COUNT_512KB) {
178- selectAlgName = "InsAllGatherConcurrentMesh1DNHR";178+ selectAlgName = "AicpuAllGatherConcurMeshNHR";
179 } else {179 } else {
180 selectAlgName = "InsAllGatherMesh1D";180 selectAlgName = "InsAllGatherMesh1D";
181 }181 }