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