Line data Source code
1 : /**
2 : * Copyright (c) 2025 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 : #include "coll_reduce_scatter_ring_zerocopy_exchange_pipeline_executor.h"
12 :
13 : namespace hccl {
14 :
15 0 : CollReduceScatterRingZerocopyExchangePipelineExecutor::CollReduceScatterRingZerocopyExchangePipelineExecutor(
16 0 : const HcclDispatcher dispatcher, std::unique_ptr<TopoMatcher> &topoMatcher)
17 0 : : CollReduceScatterExecutor(dispatcher, topoMatcher)
18 : {
19 0 : CCLMemSlice_ = false;
20 0 : DMAReduceFlag_ = true; // 设为true,以禁用RunLoop中的本地拷贝
21 0 : desc_.isZeroCopy = true; // 执行RunLoop的KernelRunInterServer分支
22 0 : desc_.deterministic = 1;
23 0 : desc_.level1SupportedAlgos = {
24 : AlgTypeLevel1::ALG_LEVEL1_RING,
25 : AlgTypeLevel1::ALG_LEVEL1_NHR,
26 : AlgTypeLevel1::ALG_LEVEL1_NB,
27 0 : };
28 0 : desc_.level2SupportedAlgos = {
29 : AlgTypeLevel2::ALG_LEVEL2_PIPELINE
30 0 : };
31 0 : }
32 :
33 0 : void CollReduceScatterRingZerocopyExchangePipelineExecutor::ParseParam(const OpParam& param)
34 : {
35 0 : tag_ = param.tag;
36 0 : root_ = param.root;
37 0 : aicpuUnfoldMode_ = param.aicpuUnfoldMode;
38 0 : opType_ = param.opType;
39 :
40 0 : u32 unitSize = SIZE_TABLE[param.DataDes.dataType];
41 0 : totalSize_ = topoAttr_.userRankSize * param.DataDes.count * unitSize;
42 0 : }
43 :
44 0 : HcclResult CollReduceScatterRingZerocopyExchangePipelineExecutor::CalcStreamNum(u32& streamNum)
45 : {
46 : // level0 需要的stream数,double ring需要2条,single ring需要1条直接用主流
47 0 : u32 totalStreamNum = (topoType_ == TopoType::TOPO_TYPE_NP_DOUBLE_RING) ?
48 : LEVEL0_PLANE_NUM_IN_NPRING_DOUBLE : 0;
49 : // level1 用NHR等ring算法,需要1条stream。但level0与level1串行,直接用主流
50 : // level2 用单ring,level2与level0/level1并行,需要1条额外的流
51 0 : totalStreamNum += 1;
52 0 : streamNum = totalStreamNum;
53 0 : HCCL_INFO("[CalcStreamNum] tag[%s] streamNum[%u] topoType_[%d]", tag_.c_str(), streamNum, topoType_);
54 :
55 0 : return HCCL_SUCCESS;
56 : }
57 :
58 0 : HcclResult CollReduceScatterRingZerocopyExchangePipelineExecutor::CalcCommInfo(
59 : std::vector<LevelNSubCommTransport>& opTransport)
60 : {
61 0 : HCCL_INFO("[CalcCommInfo] tag[%s] algoLevel0[%d] algoLevel1[%d] algoLevel2[%d]", tag_.c_str(),
62 : algType_.algoLevel0, algType_.algoLevel1, algType_.algoLevel2);
63 :
64 0 : TransportMemType inputType = TransportMemType::CCL_INPUT;
65 0 : TransportMemType outputType = TransportMemType::CCL_OUTPUT;
66 0 : CHK_RET(CalcLevel0CommInfo(inputType, outputType, opTransport));
67 0 : CHK_RET(CalcLevel1CommInfo(inputType, outputType, opTransport));
68 0 : CHK_RET(CalcLevel2CommInfo(inputType, outputType, opTransport));
69 0 : CHK_RET(CalcExchangeCommInfo(opTransport));
70 0 : return HCCL_SUCCESS;
71 : }
72 :
73 0 : HcclResult CollReduceScatterRingZerocopyExchangePipelineExecutor::CalcLevel0CommInfo(TransportMemType inputType,
74 : TransportMemType outputType, std::vector<LevelNSubCommTransport>& opTransport)
75 : {
76 0 : CommParaInfo commParaLevel0(COMM_LEVEL0, CommType::COMM_TAG_RING_INNER);
77 0 : CHK_RET(CalcCommPlaneInfo(tag_, commParaLevel0, opTransport[COMM_LEVEL0], inputType, outputType));
78 0 : LevelNSubCommTransport &commTransportLevel0 = opTransport[COMM_LEVEL0];
79 0 : for (u32 subCommIndex = 0; subCommIndex < commTransportLevel0.size(); subCommIndex++) {
80 0 : commTransportLevel0[subCommIndex].isZeroCopy = true;
81 : }
82 0 : return HCCL_SUCCESS;
83 0 : }
84 :
85 0 : HcclResult CollReduceScatterRingZerocopyExchangePipelineExecutor::CalcExchangeCommInfo(
86 : std::vector<LevelNSubCommTransport>& opTransport)
87 : {
88 0 : std::set<u32> commTargetUserRankSet;
89 0 : u32 remoteRankSend = 0;
90 0 : u32 remoteRankRecv = 0;
91 :
92 0 : CHK_RET(CalExchangeRemoteRank(remoteRankSend, remoteRankRecv));
93 0 : HCCL_INFO("[CalcExchangeCommInfo] tag[%s] userRank[%u] remoteRankSend[%u] remoteRankRecv[%u]", tag_.c_str(),
94 : topoAttr_.userRank, remoteRankSend, remoteRankRecv);
95 0 : commTargetUserRankSet.insert(remoteRankSend);
96 0 : commTargetUserRankSet.insert(remoteRankRecv);
97 : CommParaInfo commParaInfo(COMM_COMBINE_ORDER, CommType::COMM_TAG_PARTIAL_MESH_COMBINED, INVALID_VALUE_RANKID,
98 0 : INVALID_VALUE_RANKID, false, false, commTargetUserRankSet);
99 :
100 0 : TransportMemType inputType = TransportMemType::CCL_INPUT;
101 0 : TransportMemType outputType = TransportMemType::CCL_OUTPUT;
102 :
103 0 : CHK_RET(CalcCommPlaneInfo(tag_, commParaInfo, opTransport[COMM_COMBINE_ORDER], inputType, outputType));
104 0 : LevelNSubCommTransport &commTransport = opTransport[COMM_COMBINE_ORDER];
105 0 : for (u32 subCommIndex = 0; subCommIndex < commTransport.size(); subCommIndex++) {
106 0 : for (auto &transportRequest : commTransport[subCommIndex].transportRequests) {
107 0 : transportRequest.isUsedRdma = topoAttr_.isUsedRdmaMap.at(transportRequest.remoteUserRank);
108 : }
109 : }
110 0 : return HCCL_SUCCESS;
111 0 : }
112 :
113 0 : HcclResult CollReduceScatterRingZerocopyExchangePipelineExecutor::CalExchangeRemoteRank(
114 : u32 &remoteRankSend, u32 &remoteRankRecv)
115 : {
116 0 : u32 l2Size = topoAttr_.superPodNum;
117 0 : CHK_PRT_RET(l2Size == 0,
118 : HCCL_ERROR("[CalExchangeRemoteRank] invalid rank size, level2RankSize is 0"), HCCL_E_PARA);
119 0 : u32 l1Size = topoAttr_.serverNum / l2Size;
120 0 : CHK_PRT_RET(l1Size == 0,
121 : HCCL_ERROR("[CalExchangeRemoteRank] invalid rank size, level1RankSize is 0"), HCCL_E_PARA);
122 0 : u32 l0Size = topoAttr_.userRankSize / l2Size / l1Size;
123 0 : CHK_PRT_RET(l0Size == 0,
124 : HCCL_ERROR("[CalExchangeRemoteRank] invalid rank size, level0RankSize is 0"), HCCL_E_PARA);
125 :
126 : // 根据rankId计算出坐标(i, j, k)
127 0 : u32 l2Index = topoAttr_.userRank / l1Size / l0Size;
128 0 : u32 l1Index = (topoAttr_.userRank % (l1Size * l0Size)) / l0Size;
129 0 : u32 l0Index = topoAttr_.userRank % l0Size;
130 :
131 : // 计算本端将要发送数据的目标rank
132 0 : remoteRankSend = l2Index * l1Size * l0Size + l0Index * l1Size + l1Index;
133 :
134 : // 计算本端将要接收数据的目标rank
135 0 : u32 r = l1Index * l0Size + l0Index; // 超节点内相对rankid
136 0 : l0Index = r / l1Size;
137 0 : l1Index = r % l1Size;
138 0 : remoteRankRecv = l2Index * l1Size * l0Size + l1Index * l0Size + l0Index;
139 0 : return HCCL_SUCCESS;
140 : }
141 :
142 0 : u64 CollReduceScatterRingZerocopyExchangePipelineExecutor::CalcLoopMaxCount(const u32 unitSize)
143 : {
144 0 : u64 maxCountPerLoop = ((inCCLbufferSize_ / topoAttr_.serverNum / HCCL_MIN_SLICE_ALIGN) *
145 0 : HCCL_MIN_SLICE_ALIGN) / unitSize;
146 0 : return maxCountPerLoop;
147 : }
148 :
149 0 : HcclResult CollReduceScatterRingZerocopyExchangePipelineExecutor::KernelRunIntraServerPre(
150 : const OpParam ¶m, ExecMem &execMem)
151 : {
152 : (void)execMem;
153 0 : CHK_RET(SalGetDataTypeSize(param.DataDes.dataType, unitSize_));
154 0 : CHK_RET(GetCommRankInfoNormal(level0Rank_, level0RankSize_, level1Rank_, level1RankSize_, level2Rank_,
155 : level2RankSize_, false));
156 0 : CHK_RET(CalExchangeRemoteRank(exchangeRemoteRankSend_, exchangeRemoteRankRecv_));
157 :
158 0 : HCCL_INFO("[KernelRunIntraServerPre] rank[%u:%u,%u,%u], rankSize[%u, %u, %u] exchange remoteRank[send:%u Recv:%u]",
159 : topoAttr_.userRank, level2Rank_, level1Rank_, level0Rank_, level2RankSize_, level1RankSize_, level0RankSize_,
160 : exchangeRemoteRankSend_, exchangeRemoteRankRecv_);
161 0 : return HCCL_SUCCESS;
162 : }
163 :
164 0 : HcclResult CollReduceScatterRingZerocopyExchangePipelineExecutor::KernelRunInterServer(
165 : const OpParam ¶m, ExecMem &execMem)
166 : {
167 0 : curSize_ = execMem.count * unitSize_;
168 0 : HCCL_INFO("[CollReduceScatterRingZerocopyExchangePipelineExecutor] run start, rank[%u:%u,%u,%u], curSize_[%llu]",
169 : topoAttr_.userRank, level2Rank_, level1Rank_, level0Rank_, curSize_);
170 :
171 0 : for (u32 step = 0; step < level2RankSize_; step++) {
172 0 : if (!intraServerDone_) {
173 : // 只有第一个loop才需要执行节点内RS
174 0 : CHK_RET(RunIntraServer(param, execMem, step));
175 : }
176 :
177 : // 准备节点间RS的数据,user in搬运到ccl in
178 0 : CHK_RET(RunInterServerPreProcess(param, execMem, step));
179 : // 超节点内、节点间通信执行RS,编排在主流上
180 0 : if (level1RankSize_ > 1) {
181 : // 节点间RS完成后数据在ccl in
182 0 : CHK_RET(RunInterServer(param, execMem, step));
183 : }
184 : // 数据最终在ccl out
185 0 : CHK_RET(RunInterServerPostProcess(param, execMem, step));
186 :
187 : // 从steep 1开始要进行reduce,将本轮超节点间获取的数据与本轮超节点内的数据进行reduce
188 0 : if ((step > 0) && (level2RankSize_ > 1)) {
189 0 : CHK_RET(RunSuperPodPostSync(param));
190 : // 超节点间通信 与 超节点内通信 都完成后,本地进行reduce操作
191 0 : CHK_RET(RunSuperPodAndInterServerPostProcess(param, execMem, step));
192 : }
193 :
194 0 : if (step < (level2RankSize_ - 1)) {
195 : // 超节点间通信, 编排在最后一个slaveStreams上
196 0 : CHK_RET(RunSuperPodPreSync(param));
197 0 : CHK_RET(RunSuperPod(param, execMem, step + 1));
198 : }
199 : }
200 :
201 : // 将最终数据从ccl out搬到user out
202 0 : CHK_RET(RunFinallyProcess(param, execMem));
203 :
204 0 : intraServerDone_ = true;
205 0 : HCCL_INFO("[CollReduceScatterRingZerocopyExchangePipelineExecutor] run success, rank[%u:%u,%u,%u]",
206 : topoAttr_.userRank, level2Rank_, level1Rank_, level0Rank_);
207 0 : return HCCL_SUCCESS;
208 : }
209 :
210 0 : HcclResult CollReduceScatterRingZerocopyExchangePipelineExecutor::RunSuperPodPreSync(const OpParam ¶m)
211 : {
212 0 : Stream stream = param.stream;
213 0 : Stream slaveStream = algResResp_->slaveStreams.back();
214 : // 主流RS完成后,通知超节点间通信开始
215 0 : CHK_RET(LocalNotify::Post(stream, dispatcher_, algResResp_->notifiesAux.back(), INVALID_VALUE_STAGE));
216 : // 从流等待超节点内RS完成
217 0 : CHK_RET(LocalNotify::Wait(slaveStream, dispatcher_, algResResp_->notifiesAux.back(), INVALID_VALUE_STAGE));
218 0 : return HCCL_SUCCESS;
219 0 : }
220 :
221 0 : HcclResult CollReduceScatterRingZerocopyExchangePipelineExecutor::RunSuperPodPostSync(const OpParam ¶m)
222 : {
223 0 : Stream stream = param.stream;
224 0 : Stream slaveStream = algResResp_->slaveStreams.back();
225 : // 从流通知主流,超节点间数据搬运完成
226 0 : CHK_RET(LocalNotify::Post(slaveStream, dispatcher_, algResResp_->notifiesMain.back(), INVALID_VALUE_STAGE));
227 : // 主流等待超节点通信完成
228 0 : CHK_RET(LocalNotify::Wait(stream, dispatcher_, algResResp_->notifiesMain.back(), INVALID_VALUE_STAGE));
229 0 : return HCCL_SUCCESS;
230 0 : }
231 :
232 0 : HcclResult CollReduceScatterRingZerocopyExchangePipelineExecutor::RunIntraServer(
233 : const OpParam ¶m, const ExecMem &execMem, u32 step)
234 : {
235 : (void)execMem;
236 : // 计算slice信息, 将user in分成level2RankSize_块, 每个step处理一块blockIndex, 每个block需要分成level0RankSize_片
237 0 : u64 level0Count = param.DataDes.count * level1RankSize_;
238 0 : u32 blockIndex = (level2Rank_ + level2RankSize_ - (step + 1)) % level2RankSize_;
239 0 : u64 sliceSize = level0Count * unitSize_;
240 0 : u64 blockOffset = blockIndex * sliceSize * level0RankSize_;
241 :
242 0 : HCCL_DEBUG("[RunIntraServer] rank[%u:%u,%u,%u] step[%u] blockIndex[%u], level0Count[%llu]",
243 : topoAttr_.userRank, level2Rank_, level1Rank_, level0Rank_, step, blockIndex, level0Count);
244 :
245 0 : std::vector<Slice> dataSegsSlice(level0RankSize_);
246 0 : for (u32 i = 0; i < level0RankSize_; i++) {
247 0 : dataSegsSlice[i].offset = blockOffset + sliceSize * i; // 相对于param.inputPtr偏移
248 0 : dataSegsSlice[i].size = sliceSize;
249 : }
250 0 : std::vector<std::vector<Slice>> multRingsUserMemSlice = {dataSegsSlice};
251 :
252 : // 算法编排
253 0 : if (topoType_ == TopoType::TOPO_TYPE_NP_SINGLE_RING) {
254 0 : CHK_RET(MultiRingReduceScatter(param.tag, algResResp_->paramInputMem, algResResp_->paramInputMem,
255 : level0Count, param.DataDes.dataType, param.reduceType,
256 : multRingsUserMemSlice, param.stream, PROF_STAGE_1, 0, nullptr, multRingsUserMemSlice));
257 : } else {
258 0 : CHK_PRT_RET(topoType_ != TopoType::TOPO_TYPE_NP_DOUBLE_RING,
259 : HCCL_ERROR("[RunIntraServer] unknown topoType: %u", topoType_), HCCL_E_NOT_SUPPORT);
260 0 : CHK_RET(SemiRingReduceScatter(param.tag, algResResp_->paramInputMem, algResResp_->paramInputMem,
261 : level0Count, param.DataDes.dataType, param.reduceType,
262 : multRingsUserMemSlice, param.stream, PROF_STAGE_1, 0, nullptr, multRingsUserMemSlice));
263 : }
264 :
265 0 : return HCCL_SUCCESS;
266 0 : }
267 :
268 0 : HcclResult CollReduceScatterRingZerocopyExchangePipelineExecutor::SemiRingReduceScatter(
269 : const std::string &tag, DeviceMem inputMem, DeviceMem outputMem,
270 : const u64 count, const HcclDataType dataType, const HcclReduceOp reductionOp,
271 : const std::vector<std::vector<Slice> > multRingsSliceZero, Stream stream, s32 profStage,
272 : const u64 baseOffset, const HcomCollOpInfo *opInfo,
273 : const std::vector<std::vector<Slice>> multRingsUserMemSlice)
274 : {
275 : (void)tag;
276 : (void)multRingsSliceZero;
277 : (void)baseOffset;
278 : (void)opInfo;
279 0 : HCCL_DEBUG("[SemiRingReduceScatter] starts, rank[%u:%u,%u,%u]",
280 : topoAttr_.userRank, level2Rank_, level1Rank_, level0Rank_);
281 :
282 0 : CHK_RET(CheckCommSize(COMM_LEVEL0, COMM_INDEX_0 + 1));
283 0 : SubCommInfo level0CommInfo = GetSubCommInfo(COMM_LEVEL0, COMM_INDEX_0);
284 :
285 : //此处计算reduceAttr计算,outputmem使用的是scratchmem
286 0 : u64 reduceAttr = GetReduceAttr(inputMem, outputMem, dataType, reductionOp);
287 : // 执行
288 0 : std::unique_ptr<AlgTemplateBase> executor = AlgTemplateRegistry::Instance().GetAlgTemplate(
289 0 : TemplateType::TEMPLATE_REDUCESCATTER_UNIFIED_MARCH, dispatcher_);
290 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_UNIFIED_MARCH in COMM_LEVEL0", __func__);
291 0 : CHK_SMART_PTR_NULL(executor);
292 :
293 0 : CHK_RET(executor->Prepare(stream, level0CommInfo,
294 : algResResp_->paramInputMem, algResResp_->paramOutputMem, inputMem,
295 : outputMem, count, algResResp_->slaveStreams, algResResp_->notifiesMain,
296 : algResResp_->notifiesAux, dataType, reductionOp, multRingsUserMemSlice, reduceAttr));
297 :
298 0 : HcclResult ret = executor->RegisterProfiler(
299 : ((COMM_INDEX_0 + 1) << PROF_RINGINDEX_OFFSET_OF_PLANEID) +
300 0 : (level0CommInfo.localRankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID) + level0CommInfo.localRank,
301 : profStage, HCCL_EXEC_STEP_NOT_SET, stream);
302 0 : CHK_PRT_RET(ret != HCCL_SUCCESS,
303 : HCCL_ERROR("[SemiRingReduceScatter] Double ring ReduceScatter failed,return[%d]", ret), ret);
304 :
305 0 : CHK_RET(executor->RunAsync());
306 :
307 0 : HCCL_DEBUG("[SemiRingReduceScatter] run success, rank[%u:%u,%u,%u]",
308 : topoAttr_.userRank, level2Rank_, level1Rank_, level0Rank_);
309 0 : return ret;
310 0 : }
311 :
312 0 : HcclResult CollReduceScatterRingZerocopyExchangePipelineExecutor::RunInterServerPreProcess(
313 : const OpParam ¶m, const ExecMem &execMem, u32 step)
314 : {
315 : // 数据准备,将节点内RS的结果从user in搬到ccl in
316 0 : u32 blockIndex = (level2Rank_ + level2RankSize_ - (step + 1)) % level2RankSize_;
317 0 : u32 cclSliceIndex = blockIndex * level1RankSize_;
318 0 : u32 usrInSliceIndex = blockIndex * level1RankSize_ * level0RankSize_ + level1RankSize_ * level0Rank_;
319 0 : Stream stream = param.stream;
320 :
321 0 : HCCL_DEBUG("[RunInterServerPreProcess] rank[%u:%u,%u,%u] step[%u] blockIndex[%u] sliceIndex[%u, %u]",
322 : topoAttr_.userRank, level2Rank_, level1Rank_, level0Rank_, step, blockIndex, cclSliceIndex, usrInSliceIndex);
323 : // 本地 user in -> ccl in
324 0 : for (u32 i = 0; i < level1RankSize_; i++) {
325 0 : u64 ccInOffset = (cclSliceIndex + i) * curSize_;
326 0 : u64 userInOffset = (usrInSliceIndex + i) * param.DataDes.count * unitSize_; // 相对于execMem.inputPtr偏移
327 0 : DeviceMem dstMem = execMem.inputMem.range(ccInOffset, curSize_);
328 0 : DeviceMem srcMem = DeviceMem::create(static_cast<u8 *>(execMem.inputPtr) + userInOffset, curSize_);
329 0 : HCCL_DEBUG("[RunInterServerPreProcess] rank[%u:%u,%u,%u] step[%u] userInOffset[%llu] -> ccInOffset[%llu]",
330 : topoAttr_.userRank, level2Rank_, level1Rank_, level0Rank_, step, userInOffset, ccInOffset);
331 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dstMem, srcMem, stream));
332 0 : }
333 0 : return HCCL_SUCCESS;
334 0 : }
335 :
336 0 : HcclResult CollReduceScatterRingZerocopyExchangePipelineExecutor::RunInterServer(
337 : const OpParam ¶m, ExecMem &execMem, u32 step)
338 : {
339 : // 计算slice信息,也就是在ccl in的偏移
340 0 : std::vector<Slice> level1DataSegsSlice(level1RankSize_);
341 0 : u32 blockIndex = (level2Rank_ + level2RankSize_ - (step + 1)) % level2RankSize_;
342 0 : u32 sliceIndex = blockIndex * level1RankSize_;
343 :
344 0 : HCCL_DEBUG("[RunInterServer] rank[%u:%u,%u,%u] step[%u] blockIndex[%u] sliceStart[%u] sliceCnt[%u]",
345 : topoAttr_.userRank, level2Rank_, level1Rank_, level0Rank_, step, blockIndex, sliceIndex, level1RankSize_);
346 :
347 0 : for (u32 i = 0; i < level1RankSize_; i++) {
348 0 : level1DataSegsSlice[i].offset = (sliceIndex + i) * curSize_;
349 0 : level1DataSegsSlice[i].size = curSize_;
350 : }
351 :
352 0 : u64 reduceAttr = GetReduceAttr(execMem.inputMem, execMem.scratchMem, param.DataDes.dataType, param.reduceType);
353 0 : std::unique_ptr<AlgTemplateBase> level1TempAlg;
354 0 : if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_RING) {
355 : level1TempAlg =
356 0 : AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_REDUCESCATTER_RING, dispatcher_);
357 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_RING in COMM_LEVEL1", __func__);
358 0 : CHK_SMART_PTR_NULL(level1TempAlg);
359 0 : CHK_RET(level1TempAlg->Prepare(reduceAttr));
360 0 : } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NHR) {
361 : level1TempAlg =
362 0 : AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_REDUCESCATTER_NHR, dispatcher_);
363 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_NHR in COMM_LEVEL1", __func__);
364 0 : CHK_SMART_PTR_NULL(level1TempAlg);
365 0 : CHK_RET(level1TempAlg->Prepare(reduceAttr, false));
366 0 : } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NB) {
367 : level1TempAlg =
368 0 : AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_REDUCESCATTER_NB, dispatcher_);
369 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_NB in COMM_LEVEL1", __func__);
370 0 : CHK_SMART_PTR_NULL(level1TempAlg);
371 0 : CHK_RET(level1TempAlg->Prepare(reduceAttr));
372 : }
373 0 : CHK_SMART_PTR_NULL(level1TempAlg);
374 :
375 : // 执行算法编排, 主流上执行,只会使用ccl in,执行完成后数据在ccl in
376 0 : CHK_RET(CheckCommSize(COMM_LEVEL1, level0Rank_ + 1));
377 0 : SubCommInfo level1CommInfo = GetSubCommInfo(COMM_LEVEL1, level0Rank_);
378 0 : CHK_RET(level1TempAlg->Prepare(execMem.inputMem, execMem.inputMem, execMem.scratchMem, execMem.count,
379 : param.DataDes.dataType, param.stream, param.reduceType, LEVEL0_BRIDGE_RANK_ID, level1DataSegsSlice));
380 0 : CHK_RET(level1TempAlg->RegisterProfiler((level1RankSize_ << PROF_RANKSIZE_OFFSET_OF_PLANEID) + level1Rank_,
381 : PROF_STAGE_2, HCCL_EXEC_STEP_NOT_SET, param.stream));
382 0 : CHK_RET(RunTemplate(level1TempAlg, level1CommInfo));
383 :
384 0 : return HCCL_SUCCESS;
385 0 : }
386 :
387 0 : HcclResult CollReduceScatterRingZerocopyExchangePipelineExecutor::ExchangeData(
388 : const OpParam ¶m, const ExecMem &execMem, u32 step, u32 remoteRankSend, u32 remoteRankRecv)
389 : {
390 : // 获取通信对端的link
391 0 : LINK sendLink;
392 0 : LINK recvLink;
393 0 : CHK_RET(GetTransportForExchange(remoteRankSend, sendLink));
394 0 : CHK_RET(GetTransportForExchange(remoteRankRecv, recvLink));
395 0 : CHK_PTR_NULL(sendLink);
396 0 : CHK_PTR_NULL(recvLink);
397 :
398 : // 当通信对端恰好是同server的邻居时,复用Level0的建链,其注册的内存是UserMem
399 : // 否则,在CommCombineOrder上建链,其注册内存是ccl buf
400 0 : Stream stream = param.stream;
401 0 : u32 blockIndex = (level2Rank_ + level2RankSize_ - (step + 1)) % level2RankSize_;
402 0 : u32 sliceIndexSnd = blockIndex * level1RankSize_ + level1Rank_; // 要发送的数据块在本地ccl in的位置
403 0 : u32 sliceIndexCclOut = blockIndex * level1RankSize_;
404 :
405 0 : HCCL_DEBUG("[RunInterServerPostProcess] rank[%u:%u,%u,%u] step[%u] send blockIndex[%u] sliceIndex[%u] cclout[%u]",
406 : topoAttr_.userRank, level2Rank_, level1Rank_, level0Rank_, step, blockIndex, sliceIndexSnd, sliceIndexCclOut);
407 :
408 0 : bool remoteSndl0Neighbor = IsLevel0Neighbor(remoteRankSend, level0RankSize_);
409 0 : bool remoteRcvl0Neighbor = IsLevel0Neighbor(remoteRankRecv, level0RankSize_);
410 0 : if (remoteSndl0Neighbor) {
411 : // 先本地 ccl in -> user in
412 0 : u32 usrInSliceIndex = blockIndex * level1RankSize_ * level0RankSize_ +
413 0 : level1RankSize_ * level0Rank_ + level1Rank_;
414 0 : u64 userInOffset = usrInSliceIndex * param.DataDes.count * unitSize_; // 相对于param.inputPtr偏移
415 0 : DeviceMem srcMem = execMem.inputMem.range(sliceIndexSnd * curSize_, curSize_);
416 0 : DeviceMem dstMem = DeviceMem::create(static_cast<u8 *>(param.inputPtr) + userInOffset, curSize_);
417 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dstMem, srcMem, stream));
418 0 : HCCL_DEBUG("[RunInterServerPostProcess] rank[%u:%u,%u,%u] step[%u] blockIndex[%u] ci[%u]->ui[%u]",
419 : topoAttr_.userRank, level2Rank_, level1Rank_, level0Rank_, step, blockIndex, sliceIndexSnd, usrInSliceIndex);
420 :
421 : // user in send to remote user out
422 0 : CHK_RET(recvLink->TxAck(stream));
423 0 : CHK_RET(sendLink->RxAck(stream));
424 0 : CHK_RET(sendLink->TxAsync(UserMemType::OUTPUT_MEM, 0, static_cast<u8 *>(param.inputPtr) + userInOffset,
425 : curSize_, stream));
426 0 : } else {
427 : // ccl in send to remote ccl out
428 0 : CHK_RET(recvLink->TxAck(stream));
429 0 : CHK_RET(sendLink->RxAck(stream));
430 0 : CHK_RET(sendLink->TxAsync(UserMemType::OUTPUT_MEM, sliceIndexCclOut * curSize_,
431 : static_cast<u8 *>(execMem.inputMem.ptr()) + sliceIndexSnd * curSize_, curSize_, stream));
432 : }
433 :
434 0 : u32 remoteL1Rank = (remoteRankRecv % (level1RankSize_ * level0RankSize_)) / level0RankSize_;
435 0 : u32 sliceIndexRcv = blockIndex * level1RankSize_ + remoteL1Rank; // 要接收的数据块在对端ccl in的位置
436 0 : if (remoteRcvl0Neighbor) {
437 0 : u32 usrInSliceIndexPeer = blockIndex * level1RankSize_ * level0RankSize_ +
438 0 : level1Rank_ * level0RankSize_ + level0Rank_;
439 0 : u64 userInOffsetPeer = usrInSliceIndexPeer * param.DataDes.count * unitSize_; // 相对于param.inputPtr偏移
440 0 : CHK_RET(recvLink->RxAsync(UserMemType::INPUT_MEM, userInOffsetPeer, execMem.outputPtr, curSize_, stream));
441 : } else {
442 0 : CHK_RET(recvLink->RxAsync(UserMemType::INPUT_MEM, sliceIndexRcv * curSize_,
443 : static_cast<u8 *>(execMem.outputMem.ptr()) + sliceIndexCclOut * curSize_, curSize_, stream));
444 0 : CHK_RET(recvLink->PostFinAck(stream));
445 : }
446 :
447 0 : if (!remoteSndl0Neighbor) {
448 0 : CHK_RET(sendLink->WaitFinAck(stream));
449 : }
450 :
451 : // 交换数据的两端之间Barrier,确认收发完成
452 0 : CHK_RET(recvLink->TxAck(stream));
453 0 : CHK_RET(sendLink->RxAck(stream));
454 0 : CHK_RET(sendLink->TxDataSignal(stream));
455 0 : CHK_RET(recvLink->RxDataSignal(stream));
456 :
457 0 : if (remoteRcvl0Neighbor) {
458 : // 本地 user out -> ccl out
459 0 : DeviceMem srcMem = DeviceMem::create(static_cast<u8 *>(execMem.outputPtr), curSize_);
460 0 : DeviceMem dstMem = execMem.outputMem.range(sliceIndexCclOut * curSize_, curSize_);
461 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dstMem, srcMem, stream));
462 0 : HCCL_DEBUG("[RunInterServerPostProcess] rank[%u:%u,%u,%u] step[%u] blockIndex[%u] uo->co[%u]",
463 : topoAttr_.userRank, level2Rank_, level1Rank_, level0Rank_, step, blockIndex, sliceIndexCclOut);
464 0 : }
465 0 : HCCL_DEBUG("[RunInterServerPostProcess] rank[%u:%u,%u,%u] step[%u] recv blockIndex[%u] sliceIndex[%u] cclout[%u]",
466 : topoAttr_.userRank, level2Rank_, level1Rank_, level0Rank_, step, blockIndex, sliceIndexRcv, sliceIndexCclOut);
467 0 : return HCCL_SUCCESS;
468 0 : }
469 :
470 0 : HcclResult CollReduceScatterRingZerocopyExchangePipelineExecutor::RunInterServerPostProcess(
471 : const OpParam ¶m, const ExecMem &execMem, u32 step)
472 : {
473 : // 超节点内数据交换
474 0 : u32 remoteRankSend = exchangeRemoteRankSend_;
475 0 : u32 remoteRankRecv = exchangeRemoteRankRecv_;
476 :
477 0 : HCCL_DEBUG("[RunInterServerPostProcess] rank[%u:%u,%u,%u] step[%u] remoteRankSend[%u] remoteRankRecv[%u]",
478 : topoAttr_.userRank, level2Rank_, level1Rank_, level0Rank_, step, remoteRankSend, remoteRankRecv);
479 0 : if (remoteRankSend == topoAttr_.userRank && remoteRankRecv == topoAttr_.userRank) { // 不需要交换数据
480 : // 本地 ccl in -> ccl out
481 0 : Stream stream = param.stream;
482 0 : u32 blockIndex = (level2Rank_ + level2RankSize_ - (step + 1)) % level2RankSize_;
483 0 : u32 srcSliceIndex = blockIndex * level1RankSize_ + level1Rank_;
484 0 : u32 dstSliceIndex = blockIndex * level1RankSize_;
485 0 : DeviceMem srcMem = execMem.inputMem.range(srcSliceIndex * curSize_, curSize_);
486 0 : DeviceMem dstMem = execMem.outputMem.range(dstSliceIndex * curSize_, curSize_);
487 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dstMem, srcMem, stream));
488 0 : HCCL_DEBUG("[RunInterServerPostProcess] rank[%u:%u,%u,%u] step[%u] blockIndex[%u] ci[%u]->co[%u]",
489 : topoAttr_.userRank, level2Rank_, level1Rank_, level0Rank_, step, blockIndex, srcSliceIndex, dstSliceIndex);
490 0 : return HCCL_SUCCESS;
491 0 : }
492 :
493 0 : return ExchangeData(param, execMem, step, remoteRankSend, remoteRankRecv);
494 : }
495 :
496 0 : HcclResult CollReduceScatterRingZerocopyExchangePipelineExecutor::RunSuperPod(
497 : const OpParam ¶m, const ExecMem &execMem, u32 step)
498 : {
499 : (void)param;
500 0 : Stream slaveStream = algResResp_->slaveStreams.back();
501 : // 发送前回RS好的数据
502 0 : u32 blockIndexSnd = (level2Rank_ + level2RankSize_ - step) % level2RankSize_;
503 0 : u32 sliceIndexSnd = blockIndexSnd * level1RankSize_; // 要发送的数据处于本地的哪个slice
504 :
505 : // 接受上一超节点发来的数据
506 0 : u32 blockIndexRcv = (level2Rank_ + level2RankSize_ - (step + 1)) % level2RankSize_;
507 0 : u32 sliceIndexRcv = blockIndexRcv * level1RankSize_; // 要接收的数据处于对端的哪个slice
508 :
509 0 : u32 preRank = (level2Rank_ + level2RankSize_ - 1) % level2RankSize_;
510 0 : u32 nextRank = (level2Rank_ + 1) % level2RankSize_;
511 0 : CHK_RET(CheckCommSize(COMM_LEVEL2, COMM_INDEX_0 + 1));
512 0 : SubCommInfo level0CommInfo = GetSubCommInfo(COMM_LEVEL2, COMM_INDEX_0);
513 0 : LINK sendLink = level0CommInfo.links[nextRank];
514 0 : LINK recvLink = level0CommInfo.links[preRank];
515 0 : CHK_PTR_NULL(sendLink);
516 0 : CHK_PTR_NULL(recvLink);
517 :
518 : // 将数据发给nextRank前回的ccl in范围
519 0 : u32 remoteBlockIndexSndTo = (nextRank + level2RankSize_ - step) % level2RankSize_;
520 0 : u32 remoteSliceIndexSndTo = remoteBlockIndexSndTo * level1RankSize_; // 对端在哪个slice收对应的数据
521 0 : HCCL_DEBUG("[RunSuperPod] rank[%u:%u,%u,%u] step[%u] send blockIndex[%u] sliceIndex[%u]->[%u]",
522 : topoAttr_.userRank, level2Rank_, level1Rank_, level0Rank_, step, blockIndexSnd, sliceIndexSnd,
523 : remoteSliceIndexSndTo);
524 0 : HCCL_DEBUG("[RunSuperPod] rank[%u:%u,%u,%u] step[%u] recv blockIndex[%u] sliceIndex[%u]<-[%u]",
525 : topoAttr_.userRank, level2Rank_, level1Rank_, level0Rank_, step, blockIndexRcv, sliceIndexSnd, sliceIndexRcv);
526 :
527 0 : CHK_RET(recvLink->TxAck(slaveStream));
528 0 : CHK_RET(sendLink->RxAck(slaveStream));
529 : // 建链时其注册内存是ccl in与ccl out
530 : // ccl out send to remote ccl in
531 0 : CHK_RET(sendLink->TxAsync(UserMemType::INPUT_MEM, remoteSliceIndexSndTo * curSize_,
532 : static_cast<s8 *>(execMem.outputMem.ptr()) + sliceIndexSnd * curSize_, curSize_, slaveStream));
533 0 : CHK_RET(recvLink->RxAsync(UserMemType::OUTPUT_MEM, sliceIndexRcv * curSize_,
534 : static_cast<s8 *>(execMem.inputMem.ptr()) + sliceIndexSnd * curSize_, curSize_, slaveStream));
535 0 : CHK_RET(recvLink->PostFinAck(slaveStream));
536 0 : CHK_RET(sendLink->WaitFinAck(slaveStream));
537 :
538 : // 交换数据的两端之间Barrier,确认收发完成
539 0 : CHK_RET(recvLink->TxAck(slaveStream));
540 0 : CHK_RET(sendLink->RxAck(slaveStream));
541 0 : CHK_RET(sendLink->TxDataSignal(slaveStream));
542 0 : CHK_RET(recvLink->RxDataSignal(slaveStream));
543 0 : return HCCL_SUCCESS;
544 0 : }
545 :
546 0 : HcclResult CollReduceScatterRingZerocopyExchangePipelineExecutor::RunSuperPodAndInterServerPostProcess(
547 : const OpParam ¶m, const ExecMem &execMem, u32 step)
548 : {
549 : // ccl in -> ccl out执行reduce
550 0 : u32 blockIndexPreStep = (level2Rank_ + level2RankSize_ - step) % level2RankSize_;
551 0 : u32 blockIndex = (level2Rank_ + level2RankSize_ - (step + 1)) % level2RankSize_;
552 0 : u32 sliceIndex = blockIndex * level1RankSize_;
553 0 : u64 dstOffset = sliceIndex * curSize_;
554 0 : u64 srcOffset = blockIndexPreStep * level1RankSize_ * curSize_;
555 0 : HCCL_DEBUG("[RunSuperPodAndInterServerPostProcess] rank[%u:%u,%u,%u] step[%u] reduce blockIndex[%u] sliceIndex[%u]",
556 : topoAttr_.userRank, level2Rank_, level1Rank_, level0Rank_, step, blockIndex, sliceIndex);
557 :
558 0 : Stream stream = param.stream;
559 0 : CHK_RET(HcclReduceAsync(dispatcher_, static_cast<s8 *>(execMem.inputMem.ptr()) + srcOffset, execMem.count,
560 : param.DataDes.dataType, param.reduceType, stream, static_cast<s8 *>(execMem.outputMem.ptr()) + dstOffset,
561 : topoAttr_.userRank, LinkType::LINK_RESERVED, INLINE_REDUCE_BIT));
562 0 : return HCCL_SUCCESS;
563 0 : }
564 :
565 0 : HcclResult CollReduceScatterRingZerocopyExchangePipelineExecutor::RunFinallyProcess(
566 : const OpParam ¶m, const ExecMem &execMem)
567 : {
568 0 : HCCL_DEBUG("[RunFinallyProcess] rank[%u:%u,%u,%u] ccl out -> user out");
569 0 : u32 sliceIndex = level2Rank_ * level1RankSize_;
570 0 : u64 offset = sliceIndex * curSize_;
571 0 : DeviceMem srcMem = execMem.outputMem.range(offset, curSize_);
572 0 : DeviceMem dstMem = DeviceMem::create(static_cast<u8 *>(execMem.outputPtr), curSize_);
573 0 : Stream stream = param.stream;
574 0 : return HcclD2DMemcpyAsync(dispatcher_, dstMem, srcMem, stream);
575 0 : }
576 :
577 : REGISTER_EXEC("ReduceScatterRingZerocopyExchangePipelineExecutor", ReduceScatterRingZerocopyExchangePipeline,
578 : CollReduceScatterRingZerocopyExchangePipelineExecutor);
579 : } // namespace hccl
|