已开启
AicpuAllGatherConcurMeshNHR 算法适配 #5455
hulida创建于 1 天前
AicpuAllGatherConcurMeshNHR 算法适配 #5455
已开启
共 3 个文件变更+23-13
Mimpl/adv_api/detail/hccl/cc/src/ops/all_gather/executor/ins_v2_all_gather_concurrent_executor.cc+21-11
| @@ -23,7 +23,11 @@ | |||
| 23 | 23 | ||
| 24 | namespace mc2_ops_hccl { | 24 | namespace 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 | ||
| 28 | template <typename AlgTopoMatch, typename InsAlgTemplate0, typename InsAlgTemplate1> | 32 | template <typename AlgTopoMatch, typename InsAlgTemplate0, typename InsAlgTemplate1> |
| 29 | InsV2AllGatherConcurrentExecutor<AlgTopoMatch, InsAlgTemplate0, InsAlgTemplate1>::InsV2AllGatherConcurrentExecutor() | 33 | InsV2AllGatherConcurrentExecutor<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 | ||
| 160 | template <typename AlgTopoMatch, typename InsAlgTemplate0, typename InsAlgTemplate1> | 163 | template <typename AlgTopoMatch, typename InsAlgTemplate0, typename InsAlgTemplate1> |
| 161 | void InsV2AllGatherConcurrentExecutor<AlgTopoMatch, InsAlgTemplate0, InsAlgTemplate1>::GetParallelDataSplit( | 164 | void InsV2AllGatherConcurrentExecutor<AlgTopoMatch, InsAlgTemplate0, InsAlgTemplate1>::GetParallelDataSplit( |
| 162 | - std::vector<float>& splitDataSize) const | 165 | + const OpParam& param, std::vector<float>& splitDataSize) const |
| 163 | { | 166 | { |
| 164 | - const u32 portNum0 = rankSize_ - 1; // mesh端口数为rank size - 1 | 167 | + 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 | // 算法注册 |
| 391 | REGISTER_EXECUTOR_BY_TWO_TEMPS( | 401 | REGISTER_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 | 405 | ||
Mimpl/adv_api/detail/hccl/cc/src/ops/all_gather/executor/ins_v2_all_gather_concurrent_executor.h+1-1
| @@ -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 | } |