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_executor.h"
12 : #include "alg_template_register.h"
13 :
14 : namespace hccl {
15 :
16 0 : CollReduceScatterRingZerocopyExecutor::CollReduceScatterRingZerocopyExecutor(
17 0 : const HcclDispatcher dispatcher, std::unique_ptr<TopoMatcher>& topoMatcher)
18 0 : : CollReduceScatterExecutor(dispatcher, topoMatcher)
19 : {
20 0 : DMAReduceFlag_ = true; // 设为true,以禁用RunLoop中的本地拷贝
21 0 : desc_.isZeroCopy = true;
22 0 : desc_.deterministic = 1;
23 : desc_.level1SupportedAlgos
24 0 : = {AlgTypeLevel1::ALG_LEVEL1_NHR, AlgTypeLevel1::ALG_LEVEL1_NB, AlgTypeLevel1::ALG_LEVEL1_RING,
25 0 : AlgTypeLevel1::ALG_LEVEL1_AHC, AlgTypeLevel1::ALG_LEVEL1_AHC_BROKE};
26 : desc_.level2SupportedAlgos
27 0 : = {AlgTypeLevel2::ALG_LEVEL2_NHR, AlgTypeLevel2::ALG_LEVEL2_NB, AlgTypeLevel2::ALG_LEVEL2_RING};
28 0 : }
29 :
30 0 : void CollReduceScatterRingZerocopyExecutor::ParseParam(const OpParam& param)
31 : {
32 0 : tag_ = param.tag;
33 0 : totalSize_ = topoAttr_.userRankSize * param.DataDes.count * SIZE_TABLE[param.DataDes.dataType];
34 0 : aicpuUnfoldMode_ = param.aicpuUnfoldMode;
35 0 : }
36 :
37 0 : HcclResult CollReduceScatterRingZerocopyExecutor::CalcStreamNum(u32& streamNum)
38 : {
39 0 : u32 totalStreamNum = (topoType_ == TopoType::TOPO_TYPE_NP_DOUBLE_RING) ? (LEVEL0_PLANE_NUM_IN_NPRING_DOUBLE + 1) :
40 : LEVEL0_PLANE_NUM_IN_NPRING_SINGLE;
41 0 : streamNum = totalStreamNum - 1;
42 0 : HCCL_INFO("[CollReduceScatterRingZerocopyExecutor][CalcStreamNum] tag[%s] streamNum[%u]", tag_.c_str(), streamNum);
43 :
44 0 : return HCCL_SUCCESS;
45 : }
46 :
47 0 : HcclResult CollReduceScatterRingZerocopyExecutor::CalcCommInfo(std::vector<LevelNSubCommTransport>& opTransport)
48 : {
49 0 : TransportMemType inputType = TransportMemType::RESERVED;
50 0 : TransportMemType outputType = TransportMemType::RESERVED;
51 0 : CHK_RET(CalcTransportMemType(inputType, outputType));
52 0 : CHK_RET(CalcLevel0CommInfo(inputType, outputType, opTransport));
53 0 : CHK_RET(CalcLevel1CommInfo(inputType, outputType, opTransport));
54 0 : CHK_RET(CalcLevel2CommInfo(inputType, outputType, opTransport));
55 0 : return HCCL_SUCCESS;
56 : }
57 :
58 : HcclResult
59 0 : CollReduceScatterRingZerocopyExecutor::CalcTransportMemType(TransportMemType& inputType, TransportMemType& outputType)
60 : {
61 0 : inputType = TransportMemType::CCL_INPUT;
62 0 : if (scratchMemFlag_) {
63 0 : outputType = TransportMemType::SCRATCH;
64 : } else {
65 0 : outputType = TransportMemType::CCL_OUTPUT;
66 : }
67 :
68 0 : HCCL_INFO(
69 : "[CollReduceScatterRingZerocopyExecutor][CalcTransportMemType] tag[%s] inputType[%d], outputType[%d]",
70 : tag_.c_str(), inputType, outputType);
71 0 : return HCCL_SUCCESS;
72 : }
73 :
74 0 : HcclResult CollReduceScatterRingZerocopyExecutor::CalcLevel0CommInfo(
75 : TransportMemType inputType, TransportMemType outputType, std::vector<LevelNSubCommTransport>& opTransport)
76 : {
77 0 : CommParaInfo commParaLevel0(COMM_LEVEL0, CommType::COMM_TAG_RING_INNER);
78 0 : CHK_RET(CalcCommPlaneInfo(tag_, commParaLevel0, opTransport[COMM_LEVEL0], inputType, outputType));
79 0 : LevelNSubCommTransport& commTransportLevel0 = opTransport[COMM_LEVEL0];
80 0 : for (u32 subCommIndex = 0; subCommIndex < commTransportLevel0.size(); subCommIndex++) {
81 0 : commTransportLevel0[subCommIndex].isZeroCopy = true;
82 : }
83 0 : return HCCL_SUCCESS;
84 0 : }
85 :
86 0 : u64 CollReduceScatterRingZerocopyExecutor::CalcLoopMaxCount(const u32 unitSize)
87 : {
88 : // 中转内存单次最多能够接受的output count,放开ranksize限制
89 0 : u64 maxCountPerLoop
90 0 : = inCCLbufferSize_ / topoAttr_.serverNum / HCCL_MIN_SLICE_ALIGN * HCCL_MIN_SLICE_ALIGN / unitSize;
91 0 : return maxCountPerLoop;
92 : }
93 :
94 0 : HcclResult CollReduceScatterRingZerocopyExecutor::SemiRingReduceScatter(
95 : const std::string& tag, DeviceMem inputMem, DeviceMem outputMem, const u64 count, const HcclDataType dataType,
96 : const HcclReduceOp reductionOp, const std::vector<std::vector<Slice>> multRingsSliceZero, Stream stream,
97 : s32 profStage, const u64 baseOffset, const HcomCollOpInfo* opInfo,
98 : const std::vector<std::vector<Slice>> multRingsUserMemSlice)
99 : {
100 : (void)tag;
101 : (void)multRingsSliceZero;
102 : (void)baseOffset;
103 : (void)opInfo;
104 0 : HCCL_INFO("[CollReduceScatterRingZerocopyExecutor][SemiRingReduceScatter] SemiRingReduceScatter starts");
105 :
106 0 : CHK_RET(CheckCommSize(COMM_LEVEL0, COMM_INDEX_0 + 1));
107 0 : SubCommInfo level0CommInfo = GetSubCommInfo(COMM_LEVEL0, COMM_INDEX_0);
108 :
109 : // 此处计算reduceAttr计算outputmem使用scratchmem
110 0 : u64 reduceAttr = GetReduceAttr(inputMem, outputMem, dataType, reductionOp);
111 : // 执行
112 0 : std::unique_ptr<AlgTemplateBase> executor = AlgTemplateRegistry::Instance().GetAlgTemplate(
113 0 : TemplateType::TEMPLATE_REDUCESCATTER_UNIFIED_MARCH, dispatcher_);
114 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_UNIFIED_MARCH in COMM_LEVEL0", __func__);
115 0 : CHK_SMART_PTR_NULL(executor);
116 :
117 0 : CHK_RET(executor->Prepare(
118 : stream, level0CommInfo, algResResp_->paramInputMem, algResResp_->paramOutputMem, inputMem, outputMem, count,
119 : algResResp_->slaveStreams, algResResp_->notifiesMain, algResResp_->notifiesAux, dataType, reductionOp,
120 : multRingsUserMemSlice, reduceAttr));
121 0 : HCCL_DEBUG("[CollReduceScatterSemiRingExecutor][DoubleRingMidCountReduceScatter]reduceAttr is %llu", reduceAttr);
122 :
123 0 : HcclResult ret = executor->RegisterProfiler(
124 : ((COMM_INDEX_0 + 1) << PROF_RINGINDEX_OFFSET_OF_PLANEID)
125 0 : + (level0CommInfo.localRankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID) + level0CommInfo.localRank,
126 : profStage, HCCL_EXEC_STEP_NOT_SET, stream);
127 0 : CHK_PRT_RET(
128 : ret != HCCL_SUCCESS,
129 : HCCL_ERROR(
130 : "[CollReduceScatterRingZerocopyExecutor][SemiRingReduceScatter]"
131 : "Double ring ReduceScatter failed,return[%d]",
132 : ret),
133 : ret);
134 :
135 0 : CHK_RET(executor->RunAsync());
136 :
137 0 : HCCL_INFO("[CollReduceScatterRingZerocopyExecutor][SemiRingReduceScatter] SemiRingReduceScatter run success");
138 0 : return ret;
139 0 : }
140 :
141 0 : HcclResult CollReduceScatterRingZerocopyExecutor::RunIntraSeverReduceScatter(
142 : const std::string& tag, DeviceMem& inputMem, DeviceMem& outputMem, const u64 count, const HcclDataType& dataType,
143 : const HcclReduceOp& reductionOp, const std::vector<std::vector<Slice>>& multRingsSliceZero, const Stream& stream,
144 : s32 profStage, const u64 baseOffset, const HcomCollOpInfo* opInfo,
145 : const std::vector<std::vector<Slice>>& multRingsUserMemSlice, const bool disableDMAReduce)
146 : {
147 : (void)disableDMAReduce;
148 0 : if (topoType_ == TopoType::TOPO_TYPE_NP_SINGLE_RING) {
149 0 : CHK_RET(MultiRingReduceScatter(
150 : tag, inputMem, outputMem, count, dataType, reductionOp, multRingsSliceZero, stream, profStage, baseOffset,
151 : opInfo, multRingsUserMemSlice));
152 : } else {
153 0 : CHK_PRT_RET(
154 : topoType_ != TopoType::TOPO_TYPE_NP_DOUBLE_RING,
155 : HCCL_ERROR("[%s] unknown topoType: %u", __func__, topoType_), HCCL_E_NOT_SUPPORT);
156 0 : CHK_RET(SemiRingReduceScatter(
157 : tag, inputMem, outputMem, count, dataType, reductionOp, multRingsSliceZero, stream, profStage, baseOffset,
158 : opInfo, multRingsUserMemSlice));
159 : }
160 0 : return HCCL_SUCCESS;
161 : }
162 :
163 0 : HcclResult CollReduceScatterRingZerocopyExecutor::CalcLevel0DataSlices(
164 : const OpParam& param, const ExecMem& execMem, std::vector<Slice>& dataSegsSlice)
165 : {
166 0 : return CalcIntraServerDataSlicesDiscontinuous(
167 0 : param, execMem, level0RankSize_, level1RankSize_, level2RankSize_, dataSegsSlice);
168 : }
169 :
170 0 : HcclResult CollReduceScatterRingZerocopyExecutor::KernelRunIntraServerPre(const OpParam& param, ExecMem& execMem)
171 : {
172 0 : HCCL_CONFIG_INFO(
173 : HCCL_ALG,
174 : "[CollReduceScatterRingZerocopyExecutor][KernelRunIntraServerPre] The ReduceScatterDoubleRingExecutor starts");
175 0 : bool isAHCAlgo = algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC
176 0 : || algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC_BROKE;
177 0 : CHK_RET(GetCommRankInfoNormal(
178 : level0Rank_, level0RankSize_, level1Rank_, level1RankSize_, level2Rank_, level2RankSize_, isAHCAlgo));
179 :
180 : // 计算slice信息
181 0 : std::vector<Slice> dataSegsSlice;
182 0 : CHK_RET(CalcLevel0DataSlices(param, execMem, dataSegsSlice));
183 :
184 : // 执行ReduceScatter
185 0 : u64 level0Count = (dataSegsSlice.size() > level0RankSize_) ? // 如果是非连续数据通信
186 : (execMem.count) :
187 0 : (execMem.count * level1RankSize_ * level2RankSize_);
188 0 : std::vector<std::vector<Slice>> multRingsUserMemSlice = {dataSegsSlice};
189 0 : CHK_RET(RunIntraSeverReduceScatter(
190 : param.tag, execMem.inputMem, execMem.scratchMem, level0Count, param.DataDes.dataType, param.reduceType,
191 : multRingsUserMemSlice, param.stream, PROF_STAGE_1, 0, nullptr, multRingsUserMemSlice, true));
192 :
193 0 : return HCCL_SUCCESS;
194 0 : }
195 :
196 : HcclResult
197 0 : CollReduceScatterRingZerocopyExecutor::KernelRunInterServerPreProcess(const OpParam& param, const ExecMem& execMem)
198 : {
199 0 : u32 unitSize = 0;
200 0 : CHK_RET(SalGetDataTypeSize(param.DataDes.dataType, unitSize));
201 :
202 0 : DeviceMem dstMem;
203 0 : DeviceMem srcMem;
204 0 : u64 curSize = execMem.outputMem.size();
205 0 : Stream stream = param.stream;
206 0 : for (u32 i = 0; i < level1RankSize_; i++) {
207 0 : for (u32 j = 0; j < level2RankSize_; j++) {
208 : // 拷贝input上每个slice的数据到中转内存,源端每个slice的size固定为output的size
209 0 : u32 dstIndex = i * level2RankSize_ + j;
210 0 : u32 srcIndex = j * level1RankSize_ + i;
211 0 : dstMem = execMem.inputMem.range(dstIndex * curSize, curSize);
212 0 : srcMem = DeviceMem::create(
213 0 : static_cast<u8*>(execMem.inputPtr) + param.DataDes.count * unitSize * level0RankSize_ * srcIndex
214 0 : + param.DataDes.count * unitSize * level0Rank_,
215 0 : curSize);
216 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dstMem, srcMem, stream));
217 : }
218 : }
219 0 : return HCCL_SUCCESS;
220 0 : }
221 :
222 0 : HcclResult CollReduceScatterRingZerocopyExecutor::KernelRunInterServer(const OpParam& param, ExecMem& execMem)
223 : {
224 0 : u64 reduceAttr = GetReduceAttr(execMem.inputMem, execMem.scratchMem, param.DataDes.dataType, param.reduceType);
225 :
226 : // 将数据从user input搬运到ccl input
227 0 : CHK_RET(KernelRunInterServerPreProcess(param, execMem));
228 :
229 : // 计算slice
230 0 : std::vector<Slice> level1DataSegsSlice;
231 0 : CalcLevel1DataSlices(execMem.outputMem.size(), level1RankSize_, level2RankSize_, level1DataSegsSlice);
232 :
233 : // 超节点内、节点间通信
234 0 : bool isAHCAlgo = algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC
235 0 : || algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC_BROKE;
236 0 : if (level1RankSize_ > 1) {
237 : // 获取对应算法的Template
238 0 : std::unique_ptr<AlgTemplateBase> level1TempAlg;
239 0 : if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_RING) {
240 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
241 0 : TemplateType::TEMPLATE_REDUCESCATTER_RING, dispatcher_);
242 0 : CHK_SMART_PTR_NULL(level1TempAlg);
243 0 : CHK_RET(level1TempAlg->Prepare(reduceAttr));
244 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_RING in COMM_LEVEL1", __func__);
245 0 : } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NB) {
246 : level1TempAlg
247 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_REDUCESCATTER_NB, dispatcher_);
248 0 : CHK_SMART_PTR_NULL(level1TempAlg);
249 0 : CHK_RET(level1TempAlg->Prepare(reduceAttr));
250 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_NB in COMM_LEVEL1", __func__);
251 0 : } else if (isAHCAlgo) {
252 : // 获取通信域分组信息
253 0 : std::vector<std::vector<std::vector<u32>>> globalSubGroups;
254 0 : std::map<AHCConcOpType, TemplateType> ahcAlgOption;
255 0 : CHK_RET(topoMatcher_->GetGlobalSubGroups(COMM_LEVEL1_AHC, globalSubGroups));
256 0 : topoMatcher_->GetAHCAlgOption(ahcAlgOption);
257 0 : if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC) {
258 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
259 0 : TemplateType::TEMPLATE_REDUCESCATTER_AHC, dispatcher_);
260 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_AHC in COMM_LEVEL1", __func__);
261 : } else {
262 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
263 0 : TemplateType::TEMPLATE_REDUCESCATTER_AHC_BROKE, dispatcher_);
264 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_AHC_BROKE in COMM_LEVEL1", __func__);
265 : }
266 0 : CHK_SMART_PTR_NULL(level1TempAlg);
267 0 : CHK_RET(level1TempAlg->Prepare(execMem.count, globalSubGroups, ahcAlgOption));
268 0 : CHK_RET(level1TempAlg->Prepare(reduceAttr));
269 0 : } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NHR) {
270 : level1TempAlg
271 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_REDUCESCATTER_NHR, dispatcher_);
272 0 : CHK_SMART_PTR_NULL(level1TempAlg);
273 0 : CHK_RET(level1TempAlg->Prepare(reduceAttr, false));
274 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_NHR in COMM_LEVEL1", __func__);
275 : } else {
276 0 : HCCL_ERROR("ReduceScatter ring: unsupported level1 algtype [%s]", AlgTypeToStr(algType_).c_str());
277 0 : return HCCL_E_NOT_SUPPORT;
278 : }
279 : // 执行算法编排
280 0 : CommPlane commPlaneLevel1 = isAHCAlgo ? COMM_LEVEL1_AHC : COMM_LEVEL1;
281 0 : CHK_RET(CheckCommSize(commPlaneLevel1, level0Rank_ + 1));
282 0 : SubCommInfo level1CommInfo = GetSubCommInfo(commPlaneLevel1, level0Rank_);
283 0 : CHK_RET(level1TempAlg->Prepare(
284 : execMem.inputMem, execMem.inputMem, execMem.scratchMem, execMem.count, param.DataDes.dataType, param.stream,
285 : param.reduceType, LEVEL0_BRIDGE_RANK_ID, level1DataSegsSlice));
286 0 : CHK_RET(level1TempAlg->RegisterProfiler(
287 : (level1RankSize_ << PROF_RANKSIZE_OFFSET_OF_PLANEID) + level1Rank_, PROF_STAGE_2, HCCL_EXEC_STEP_NOT_SET,
288 : param.stream));
289 0 : CHK_RET(RunTemplate(level1TempAlg, level1CommInfo));
290 0 : }
291 :
292 : // 超节点间通信
293 0 : if (level2RankSize_ > 1 && !isAHCAlgo) {
294 : // 获取对应算法的Template
295 0 : std::unique_ptr<AlgTemplateBase> level2TempAlg;
296 0 : if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NB) {
297 : level2TempAlg
298 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_REDUCESCATTER_NB, dispatcher_);
299 0 : CHK_SMART_PTR_NULL(level2TempAlg);
300 0 : CHK_RET(level2TempAlg->Prepare(reduceAttr));
301 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_NB in COMM_LEVEL2", __func__);
302 0 : } else if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_RING) {
303 0 : level2TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
304 0 : TemplateType::TEMPLATE_REDUCESCATTER_RING, dispatcher_);
305 0 : CHK_SMART_PTR_NULL(level2TempAlg);
306 0 : CHK_RET(level2TempAlg->Prepare(reduceAttr));
307 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_RING in COMM_LEVEL2", __func__);
308 0 : } else if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NHR) {
309 : level2TempAlg
310 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_REDUCESCATTER_NHR, dispatcher_);
311 0 : CHK_SMART_PTR_NULL(level2TempAlg);
312 0 : CHK_RET(level2TempAlg->Prepare(reduceAttr, false));
313 0 : if (algoAttr_.isSupportAtomicWrite) {
314 0 : level2TempAlg->CloseBarrier();
315 : }
316 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_NHR in COMM_LEVEL2", __func__);
317 : } else {
318 0 : HCCL_ERROR("ReduceScatter ring: unsupported level2 algtype [%s]", AlgTypeToStr(algType_).c_str());
319 0 : return HCCL_E_NOT_SUPPORT;
320 : }
321 : // 执行算法编排
322 : DeviceMem level2InputMem
323 0 : = execMem.inputMem.range(level1DataSegsSlice[level1Rank_].offset, level1DataSegsSlice[level1Rank_].size);
324 0 : CHK_RET(level2TempAlg->Prepare(
325 : level2InputMem, level2InputMem, execMem.scratchMem, execMem.count, param.DataDes.dataType, param.stream,
326 : param.reduceType, LEVEL0_BRIDGE_RANK_ID, std::vector<Slice>(0), level1DataSegsSlice[level1Rank_].offset));
327 0 : CHK_RET(level2TempAlg->RegisterProfiler(
328 : (level2RankSize_ << PROF_RANKSIZE_OFFSET_OF_PLANEID) + level2Rank_, PROF_STAGE_2, HCCL_EXEC_STEP_NOT_SET,
329 : param.stream));
330 0 : CHK_RET(CheckCommSize(COMM_LEVEL2, COMM_INDEX_0 + 1));
331 0 : SubCommInfo level2CommInfo = GetSubCommInfo(COMM_LEVEL2, COMM_INDEX_0);
332 0 : CHK_RET(RunTemplate(level2TempAlg, level2CommInfo));
333 0 : }
334 :
335 : // 后处理
336 0 : CHK_RET(KernelRunInterServerPostProcess(param, execMem));
337 :
338 0 : return HCCL_SUCCESS;
339 0 : }
340 :
341 : HcclResult
342 0 : CollReduceScatterRingZerocopyExecutor::KernelRunInterServerPostProcess(const OpParam& param, const ExecMem& execMem)
343 : {
344 0 : u32 dataIndex = level1Rank_ * level2RankSize_ + level2Rank_;
345 0 : u64 curSize = execMem.outputMem.size();
346 0 : DeviceMem srcMem = execMem.inputMem.range(curSize * dataIndex, curSize);
347 0 : DeviceMem dstMem = DeviceMem::create(static_cast<u8*>(execMem.outputPtr), curSize);
348 0 : Stream stream = param.stream;
349 0 : return HcclD2DMemcpyAsync(dispatcher_, dstMem, srcMem, stream);
350 0 : }
351 :
352 : REGISTER_EXEC("ReduceScatterRingZerocopyExecutor", ReduceScatterRingZerocopy, CollReduceScatterRingZerocopyExecutor);
353 : } // namespace hccl
|