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_batch_send_recv_group_executor.h"
12 :
13 : namespace hccl {
14 :
15 64 : CollBatchSendRecvGroupExecutor::CollBatchSendRecvGroupExecutor(
16 64 : const HcclDispatcher dispatcher, std::unique_ptr<TopoMatcher>& topoMatcher)
17 64 : : CollBatchSendRecvExecutor(dispatcher, topoMatcher)
18 64 : {}
19 :
20 2 : void CollBatchSendRecvGroupExecutor::CalcIsSuPodAsym(bool isA2MultiModule)
21 : {
22 : // isSuPodAsym 表示A2A3卡数不一致场景或者A3多超节点server数不同场景
23 : // 参考alltoallv_direct_fullmesh executor中的判断逻辑
24 2 : if (topoAttr_.superPodNum > 1) {
25 : isSuPodAsym_
26 0 : = (static_cast<bool>(topoAttr_.multiModuleDiffDeviceNumMode)
27 0 : || static_cast<bool>(topoAttr_.multiSuperPodDiffServerNumMode));
28 : } else {
29 6 : isSuPodAsym_ = (static_cast<bool>(topoMatcher_->GetExternalInputInterHccsDisable()) || isA2MultiModule)
30 4 : && static_cast<bool>(topoAttr_.multiModuleDiffDeviceNumMode);
31 : }
32 2 : HCCL_INFO(
33 : "[CalcIsSuPodAsym] superPodNum[%u] isA2MultiModule[%d] isSuPodAsym[%d]", topoAttr_.superPodNum, isA2MultiModule,
34 : isSuPodAsym_);
35 2 : }
36 :
37 2 : HcclResult CollBatchSendRecvGroupExecutor::CalcPingPongHalfSize()
38 : {
39 2 : u32 pingPongSliceNum = GROUP_MAX_CONCURRENT * 2;
40 2 : u32 alignSize = HCCL_MIN_SLICE_ALIGN_910B;
41 2 : bufferSliceSize_ = algResResp_->cclInputMem.size() / alignSize / pingPongSliceNum * alignSize;
42 : // RDMA单流slot大小 = CCLOut / 2:A半区(send, 单流单slot)、B半区(recv, 单流单slot)
43 2 : rdmaDataBlockSize_ = algResResp_->cclOutputMem.size() / alignSize / RDMA_CCLOUT_HALF_NUM * alignSize;
44 2 : HCCL_INFO(
45 : "[CollBatchSendRecvGroupExecutor][CalcPingPongHalfSize] pingPong halfSize[%llu] rdmaDataBlockSize_[%llu]",
46 : bufferSliceSize_, rdmaDataBlockSize_);
47 2 : return HCCL_SUCCESS;
48 : }
49 :
50 1 : HcclResult CollBatchSendRecvGroupExecutor::OrganizeSendItemByStream()
51 : {
52 1 : sendQueueBySendstream_.resize(sendStreamNum_);
53 1 : HCCL_INFO("[OrganizeSendItemByStream] sendStreamNum_[%u]", sendStreamNum_);
54 3 : while (!sendDeque_.empty()) {
55 2 : HcclSendRecvItem* curr = sendDeque_.front();
56 2 : CHK_PTR_NULL(curr);
57 2 : sendQueueBySendstream_[curr->remoteRank % sendStreamNum_].push_back(curr);
58 2 : sendDeque_.pop_front();
59 : }
60 1 : HCCL_INFO("OrganizeSendItemByStream Done!");
61 1 : return HCCL_SUCCESS;
62 : }
63 :
64 2 : HcclResult CollBatchSendRecvGroupExecutor::OrganizeRecvItemByStream()
65 : {
66 2 : recvQueueByRecvstream_.resize(recvStreamNum_);
67 2 : HCCL_INFO("[OrganizeRecvItemByStream] recvStreamNum_[%u]", recvStreamNum_);
68 4 : while (!recvDeque_.empty()) {
69 2 : HcclSendRecvItem* curr = recvDeque_.front();
70 2 : CHK_PTR_NULL(curr);
71 2 : recvQueueByRecvstream_[curr->remoteRank % recvStreamNum_].push_back(curr);
72 2 : recvDeque_.pop_front();
73 : }
74 2 : HCCL_INFO("OrganizeRecvItemByStream Done!");
75 2 : return HCCL_SUCCESS;
76 : }
77 :
78 2 : HcclResult CollBatchSendRecvGroupExecutor::CalcPodRange()
79 : {
80 : // Determine pod (server/supernode) membership — same approach as alltoallv_direct_fullmesh
81 2 : u32 devNumInlocalPod = INVALID_VALUE_RANKSIZE;
82 2 : u32 rankIdxInPod = INVALID_VALUE_RANKID;
83 2 : bool isA2MultiModule = topoAttr_.deviceType == DevType::DEV_TYPE_910B && !topoAttr_.isSingleMeshAggregation;
84 2 : if (static_cast<bool>(topoMatcher_->GetExternalInputInterHccsDisable()) || isA2MultiModule) {
85 1 : CHK_RET(topoMatcher_->GetLocalServerRankSize(topoAttr_.userRank, devNumInlocalPod, rankIdxInPod));
86 : } else {
87 1 : CHK_RET(topoMatcher_->GetLocalSuperPodRankSize(topoAttr_.userRank, devNumInlocalPod, rankIdxInPod));
88 : }
89 2 : podStartRank_ = topoAttr_.userRank - rankIdxInPod;
90 2 : podEndRank_ = podStartRank_ + devNumInlocalPod - 1;
91 2 : devNumInlocalPod_ = devNumInlocalPod;
92 2 : CalcIsSuPodAsym(isA2MultiModule);
93 2 : HCCL_INFO(
94 : "[CalcPodRange] userRank[%u] pod[%u-%u] devNumInlocalPod[%u] rankIdxInPod[%u] isSuPodAsym[%d]",
95 : topoAttr_.userRank, podStartRank_, podEndRank_, devNumInlocalPod, rankIdxInPod, isSuPodAsym_);
96 2 : return HCCL_SUCCESS;
97 : }
98 :
99 35 : bool CollBatchSendRecvGroupExecutor::IsRemoteRankRdma(u32 remoteRank) const
100 : {
101 : // pod内为SDMA,跨pod为RDMA
102 35 : return !(remoteRank >= podStartRank_ && remoteRank <= podEndRank_);
103 : }
104 :
105 0 : HcclResult CollBatchSendRecvGroupExecutor::CalcCommInfo(std::vector<LevelNSubCommTransport>& opTransport)
106 : {
107 0 : CHK_RET(CollBatchSendRecvExecutor::CalcCommInfo(opTransport));
108 :
109 0 : LevelNSubCommTransport& commTransport = opTransport[COMM_COMBINE_ORDER];
110 0 : for (u32 subCommIndex = 0; subCommIndex < commTransport.size(); subCommIndex++) {
111 0 : for (auto& transportRequest : commTransport[subCommIndex].transportRequests) {
112 0 : transportRequest.isUsedRdma = topoAttr_.isUsedRdmaMap.at(transportRequest.remoteUserRank);
113 : }
114 : }
115 0 : return HCCL_SUCCESS;
116 : }
117 :
118 1 : HcclResult CollBatchSendRecvGroupExecutor::Orchestrate(OpParam& param, AlgResourceResponse& algResource)
119 : {
120 1 : HcclUs startut = TIME_NOW();
121 1 : HCCL_CONFIG_INFO(HCCL_ALG, "[CollBatchSendRecvGroupExecutor] groupsendrecv starts.");
122 :
123 1 : sendStreamNum_ = GROUP_MAX_CONCURRENT;
124 1 : recvStreamNum_ = GROUP_MAX_CONCURRENT;
125 1 : HCCL_INFO("[Orchestrate] sendStreamNum_[%u], recvStreamNum_[%u]", sendStreamNum_, recvStreamNum_);
126 :
127 1 : algResResp_ = &algResource;
128 1 : CHK_RET(CheckCommSize(COMM_COMBINE_ORDER, COMM_SIZE_TWO));
129 1 : CHK_RET(GetPairWiseList(param.BatchSendRecvDataDes.sendRecvItemsPtr, param.BatchSendRecvDataDes.itemNum));
130 0 : CHK_RET(ProcessSelfSendRecvTasks(param.stream));
131 0 : if (topoAttr_.userRankSize == 1) {
132 0 : HCCL_INFO(
133 : "tag[%s] BatchSendRecvGroup Executor orchestrate success, take time [%lld]us.", param.tag.c_str(),
134 : DURATION_US(TIME_NOW() - startut));
135 0 : return HCCL_SUCCESS;
136 : }
137 0 : CHK_RET(CalcPodRange());
138 0 : CHK_RET(CalcPingPongHalfSize()); // ping-pong double buffering + RDMA slot
139 0 : CHK_RET(OrganizeSendItemByStream());
140 0 : CHK_RET(OrganizeRecvItemByStream());
141 :
142 0 : CHK_RET(CalcSendSlices());
143 0 : CHK_RET(CalcRecvSlices());
144 :
145 0 : CHK_RET(RunLoop(param));
146 :
147 0 : HCCL_INFO(
148 : "tag[%s] BatchSendRecvGroup Executor orchestrate success, take time [%lld]us.", param.tag.c_str(),
149 : DURATION_US(TIME_NOW() - startut));
150 0 : return HCCL_SUCCESS;
151 : }
152 :
153 4 : HcclResult CollBatchSendRecvGroupExecutor::CalcStreamTaskStatus(u32& nonEmptySendStream, u32& nonEmptyRecvStream)
154 : {
155 4 : nonEmptySendStream = 0;
156 4 : nonEmptyRecvStream = 0;
157 : // 记录各从流是否有任务,供头尾同步只唤醒有任务的从流(循环中不再更新)。
158 4 : sendStreamHasTask_.assign(sendStreamNum_, false);
159 4 : recvStreamHasTask_.assign(recvStreamNum_, false);
160 19 : for (u32 i = 0; i < sendStreamNum_; i++) {
161 15 : if (!sendDataSlicesBySendStream_[i].empty()) {
162 5 : nonEmptySendStream++;
163 5 : sendStreamHasTask_[i] = true;
164 : }
165 : }
166 18 : for (u32 i = 0; i < recvStreamNum_; i++) {
167 14 : if (!recvDataSlicesByRecvStream_[i].empty()) {
168 2 : nonEmptyRecvStream++;
169 2 : recvStreamHasTask_[i] = true;
170 : }
171 : }
172 4 : rdmaSendHasTask_ = !rdmaSendSlices_.empty();
173 4 : rdmaRecvHasTask_ = !rdmaRecvSlices_.empty();
174 4 : return HCCL_SUCCESS;
175 : }
176 :
177 1 : HcclResult CollBatchSendRecvGroupExecutor::MainPostSubWait(Stream& mainStream)
178 : {
179 : // 主流只通知有任务的从流开始
180 3 : for (u32 i = 0; i < sendStreamNum_; i++) {
181 2 : if (!sendStreamHasTask_[i]) {
182 2 : continue;
183 : }
184 0 : CHK_RET(LocalNotify::Post(mainStream, dispatcher_, algResResp_->notifiesAux[i], PROF_STAGE_0));
185 0 : CHK_RET(
186 : LocalNotify::Wait(algResResp_->slaveStreams[i], dispatcher_, algResResp_->notifiesAux[i], PROF_STAGE_0));
187 0 : HCCL_DEBUG("MainPost, Send[%u] Wait", i);
188 : }
189 :
190 3 : for (u32 i = 0; i < recvStreamNum_; i++) {
191 2 : if (!recvStreamHasTask_[i]) {
192 2 : continue;
193 : }
194 0 : CHK_RET(LocalNotify::Post(mainStream, dispatcher_, algResResp_->notifiesAux[i + sendStreamNum_], PROF_STAGE_0));
195 0 : CHK_RET(LocalNotify::Wait(
196 : algResResp_->slaveStreams[i + sendStreamNum_], dispatcher_, algResResp_->notifiesAux[i + sendStreamNum_],
197 : PROF_STAGE_0));
198 0 : HCCL_DEBUG("MainPost, Recv[%u] Wait", i);
199 : }
200 :
201 : // RDMA专用从流(send/recv各一条)
202 1 : u32 rdmaSendIdx = RdmaSendStreamIdx();
203 1 : u32 rdmaRecvIdx = RdmaRecvStreamIdx();
204 1 : if (rdmaSendHasTask_) {
205 0 : CHK_RET(LocalNotify::Post(mainStream, dispatcher_, algResResp_->notifiesAux[rdmaSendIdx], PROF_STAGE_0));
206 0 : CHK_RET(LocalNotify::Wait(
207 : algResResp_->slaveStreams[rdmaSendIdx], dispatcher_, algResResp_->notifiesAux[rdmaSendIdx], PROF_STAGE_0));
208 0 : HCCL_DEBUG("MainPost, RdmaSend Wait");
209 : }
210 1 : if (rdmaRecvHasTask_) {
211 0 : CHK_RET(LocalNotify::Post(mainStream, dispatcher_, algResResp_->notifiesAux[rdmaRecvIdx], PROF_STAGE_0));
212 0 : CHK_RET(LocalNotify::Wait(
213 : algResResp_->slaveStreams[rdmaRecvIdx], dispatcher_, algResResp_->notifiesAux[rdmaRecvIdx], PROF_STAGE_0));
214 0 : HCCL_DEBUG("MainPost, RdmaRecv Wait");
215 : }
216 :
217 1 : return HCCL_SUCCESS;
218 : }
219 :
220 1 : HcclResult CollBatchSendRecvGroupExecutor::MainWaitSubPost(Stream& mainStream)
221 : {
222 : // 最后主流只等待有任务的从流结束
223 3 : for (u32 i = 0; i < sendStreamNum_; i++) {
224 2 : if (!sendStreamHasTask_[i]) {
225 2 : continue;
226 : }
227 0 : CHK_RET(
228 : LocalNotify::Post(algResResp_->slaveStreams[i], dispatcher_, algResResp_->notifiesMain[i], PROF_STAGE_0));
229 0 : CHK_RET(LocalNotify::Wait(mainStream, dispatcher_, algResResp_->notifiesMain[i], PROF_STAGE_0));
230 0 : HCCL_DEBUG("MainWait, Send[%u] Post", i);
231 : }
232 :
233 3 : for (u32 i = 0; i < recvStreamNum_; i++) {
234 2 : if (!recvStreamHasTask_[i]) {
235 2 : continue;
236 : }
237 0 : CHK_RET(LocalNotify::Post(
238 : algResResp_->slaveStreams[i + sendStreamNum_], dispatcher_, algResResp_->notifiesMain[i + sendStreamNum_],
239 : PROF_STAGE_0));
240 0 : CHK_RET(
241 : LocalNotify::Wait(mainStream, dispatcher_, algResResp_->notifiesMain[i + sendStreamNum_], PROF_STAGE_0));
242 0 : HCCL_DEBUG("MainWait, Recv[%u] Post", i);
243 : }
244 :
245 : // RDMA专用从流(send/recv各一条)
246 1 : u32 rdmaSendIdx = RdmaSendStreamIdx();
247 1 : u32 rdmaRecvIdx = RdmaRecvStreamIdx();
248 1 : if (rdmaSendHasTask_) {
249 0 : CHK_RET(LocalNotify::Post(
250 : algResResp_->slaveStreams[rdmaSendIdx], dispatcher_, algResResp_->notifiesMain[rdmaSendIdx], PROF_STAGE_0));
251 0 : CHK_RET(LocalNotify::Wait(mainStream, dispatcher_, algResResp_->notifiesMain[rdmaSendIdx], PROF_STAGE_0));
252 0 : HCCL_DEBUG("MainWait, RdmaSend Post");
253 : }
254 1 : if (rdmaRecvHasTask_) {
255 0 : CHK_RET(LocalNotify::Post(
256 : algResResp_->slaveStreams[rdmaRecvIdx], dispatcher_, algResResp_->notifiesMain[rdmaRecvIdx], PROF_STAGE_0));
257 0 : CHK_RET(LocalNotify::Wait(mainStream, dispatcher_, algResResp_->notifiesMain[rdmaRecvIdx], PROF_STAGE_0));
258 0 : HCCL_DEBUG("MainWait, RdmaRecv Post");
259 : }
260 :
261 1 : return HCCL_SUCCESS;
262 : }
263 :
264 : HcclResult
265 2 : CollBatchSendRecvGroupExecutor::ProcessPreloadedSendSlice(u32 streamIdx, u32& pendingSendCount, u32& nonEmptySendStream)
266 : {
267 2 : u32 curPhase = sendCurPhase_[streamIdx];
268 2 : u32 loadedRank = sendLoadedRemoteRank_[streamIdx];
269 2 : u64 loadedSize = sendLoadedSize_[streamIdx];
270 2 : u64 kernelOffset = bufferSliceSize_ * (streamIdx * 2 + curPhase);
271 2 : u64 phaseOffset = bufferSliceSize_ * curPhase;
272 :
273 2 : HCCL_INFO(
274 : "[RunTasks] SendStream[%u](loaded) phase[%u] offset[%llu] size[%llu] rank[%u]", streamIdx, curPhase,
275 : kernelOffset, loadedSize, loadedRank);
276 :
277 : // Step A.1: Record — notify remote data ready (TxPrepare + TxData)
278 2 : LINK sendTargetLink;
279 2 : CHK_RET(GetSendTargetLink(loadedRank, sendTargetLink));
280 2 : DeviceMem inCommMem = algResResp_->cclInputMem.range(kernelOffset, loadedSize);
281 2 : CHK_RET(sendTargetLink->TxPrepare(algResResp_->slaveStreams[streamIdx]));
282 2 : CHK_RET(sendTargetLink->TxData(
283 : UserMemType::OUTPUT_MEM, phaseOffset, inCommMem.ptr(), loadedSize, algResResp_->slaveStreams[streamIdx]));
284 :
285 2 : sendLoadedSize_[streamIdx] = 0;
286 2 : pendingSendCount--;
287 :
288 : // Step A.2: D2D next slice to OTHER half (only if same rank)
289 2 : if (!sendDataSlicesBySendStream_[streamIdx].empty()) {
290 1 : SendRecvSlice& nextSlice = sendDataSlicesBySendStream_[streamIdx].front();
291 1 : if (nextSlice.remoteRank == loadedRank) {
292 1 : u32 nextPhase = 1 - curPhase;
293 1 : u64 d2dOffset = bufferSliceSize_ * (streamIdx * 2 + nextPhase);
294 1 : DeviceMem d2dCommMem = algResResp_->cclInputMem.range(d2dOffset, nextSlice.size);
295 1 : DeviceMem inMem(nextSlice.addr, nextSlice.size);
296 1 : HCCL_INFO(
297 : "[RunTasks] SendStream[%u] load next to phase[%u] offset[%llu] size[%llu]", streamIdx, nextPhase,
298 : d2dOffset, nextSlice.size);
299 1 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, d2dCommMem, inMem, algResResp_->slaveStreams[streamIdx]));
300 1 : sendLoadedSize_[streamIdx] = nextSlice.size;
301 1 : sendLoadedRemoteRank_[streamIdx] = nextSlice.remoteRank;
302 1 : sendCurPhase_[streamIdx] = nextPhase;
303 1 : sendDataSlicesBySendStream_[streamIdx].pop_front();
304 1 : pendingSendCount++;
305 1 : if (sendDataSlicesBySendStream_[streamIdx].empty()) {
306 1 : nonEmptySendStream--;
307 1 : HCCL_INFO("[RunTasks] nonEmptySendStream[%u]", nonEmptySendStream);
308 : }
309 1 : }
310 : }
311 :
312 : // Step A.3: Wait for remote ack (TxDone)
313 2 : CHK_RET(sendTargetLink->TxDone(algResResp_->slaveStreams[streamIdx]));
314 2 : return HCCL_SUCCESS;
315 2 : }
316 :
317 : HcclResult
318 2 : CollBatchSendRecvGroupExecutor::ProcessNewRankSendSlice(u32 streamIdx, u32& pendingSendCount, u32& nonEmptySendStream)
319 : {
320 2 : SendRecvSlice& firstSlice = sendDataSlicesBySendStream_[streamIdx].front();
321 2 : u32 newRank = firstSlice.remoteRank;
322 :
323 : // Step B.1: D2D first slice to half A (phase 0)
324 2 : u64 offsetA = bufferSliceSize_ * (streamIdx * 2 + 0);
325 2 : DeviceMem commMemA = algResResp_->cclInputMem.range(offsetA, firstSlice.size);
326 2 : DeviceMem inMem(firstSlice.addr, firstSlice.size);
327 2 : HCCL_INFO(
328 : "[RunTasks] SendStream[%u] preload rank[%u] to A offset[%llu] size[%llu]", streamIdx, newRank, offsetA,
329 : firstSlice.size);
330 2 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, commMemA, inMem, algResResp_->slaveStreams[streamIdx]));
331 :
332 2 : sendDataSlicesBySendStream_[streamIdx].pop_front();
333 :
334 : // Step B.2: Record from A
335 2 : LINK sendTargetLink;
336 2 : CHK_RET(GetSendTargetLink(newRank, sendTargetLink));
337 2 : CHK_RET(sendTargetLink->TxPrepare(algResResp_->slaveStreams[streamIdx]));
338 2 : CHK_RET(sendTargetLink->TxData(
339 : UserMemType::OUTPUT_MEM, 0, commMemA.ptr(), firstSlice.size, algResResp_->slaveStreams[streamIdx]));
340 :
341 : // Step B.3: D2D next to B (only if same rank)
342 2 : sendLoadedSize_[streamIdx] = 0;
343 2 : sendLoadedRemoteRank_[streamIdx] = newRank;
344 2 : sendCurPhase_[streamIdx] = 0;
345 2 : if (!sendDataSlicesBySendStream_[streamIdx].empty()) {
346 1 : SendRecvSlice& nextSlice = sendDataSlicesBySendStream_[streamIdx].front();
347 1 : if (nextSlice.remoteRank == newRank) {
348 1 : u64 offsetB = bufferSliceSize_ * (streamIdx * 2 + 1);
349 1 : DeviceMem commMemB = algResResp_->cclInputMem.range(offsetB, nextSlice.size);
350 1 : DeviceMem inMemB(nextSlice.addr, nextSlice.size);
351 1 : HCCL_INFO(
352 : "[RunTasks] SendStream[%u] load next to B offset[%llu] size[%llu]", streamIdx, offsetB, nextSlice.size);
353 1 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, commMemB, inMemB, algResResp_->slaveStreams[streamIdx]));
354 1 : sendLoadedSize_[streamIdx] = nextSlice.size;
355 1 : sendLoadedRemoteRank_[streamIdx] = nextSlice.remoteRank;
356 1 : sendCurPhase_[streamIdx] = 1;
357 1 : sendDataSlicesBySendStream_[streamIdx].pop_front();
358 1 : pendingSendCount++;
359 1 : if (sendDataSlicesBySendStream_[streamIdx].empty()) {
360 1 : nonEmptySendStream--;
361 1 : HCCL_INFO("[RunTasks] nonEmptySendStream[%u]", nonEmptySendStream);
362 : }
363 1 : }
364 : } else {
365 1 : nonEmptySendStream--;
366 1 : HCCL_INFO("[RunTasks] nonEmptySendStream[%u]", nonEmptySendStream);
367 : }
368 :
369 : // Step B.4: Wait for remote ack (TxDone)
370 2 : CHK_RET(sendTargetLink->TxDone(algResResp_->slaveStreams[streamIdx]));
371 2 : return HCCL_SUCCESS;
372 2 : }
373 :
374 2 : HcclResult CollBatchSendRecvGroupExecutor::ProcessRecvSlice(u32 streamIdx, u32& nonEmptyRecvStream)
375 : {
376 2 : SendRecvSlice& slice = recvDataSlicesByRecvStream_[streamIdx].front();
377 :
378 : // Reset phase to 0 when encountering a new rank
379 2 : if (slice.remoteRank != recvCurRemoteRank_[streamIdx]) {
380 2 : recvCurPhase_[streamIdx] = 0;
381 2 : recvCurRemoteRank_[streamIdx] = slice.remoteRank;
382 : }
383 :
384 2 : u32 curPhase = recvCurPhase_[streamIdx];
385 2 : u64 offset = bufferSliceSize_ * (streamIdx * 2 + curPhase);
386 2 : u64 phaseOffset = bufferSliceSize_ * (curPhase + (topoAttr_.userRank % sendStreamNum_) * 2);
387 2 : HCCL_INFO(
388 : "[RunTasks] RecvStream[%u] phase[%u] offset[%llu] phaseOffset[%llu] size[%llu] rank[%u]", streamIdx, curPhase,
389 : offset, phaseOffset, slice.size, slice.remoteRank);
390 :
391 2 : LINK recvTargetLink;
392 2 : CHK_RET(GetRecvTargetLink(slice.remoteRank, recvTargetLink));
393 :
394 2 : CHK_RET(recvTargetLink->RxPrepare(algResResp_->slaveStreams[streamIdx + sendStreamNum_]));
395 2 : DeviceMem outMem(slice.addr, slice.size);
396 2 : CHK_RET(recvTargetLink->RxData(
397 : UserMemType::INPUT_MEM, phaseOffset, outMem.ptr(), slice.size,
398 : algResResp_->slaveStreams[streamIdx + sendStreamNum_]));
399 2 : HCCL_INFO("[RunTasks] RecvStream[%u] direct, outMem ptr[%p], size[%llu]", streamIdx, outMem.ptr(), outMem.size());
400 2 : CHK_RET(recvTargetLink->RxDone(algResResp_->slaveStreams[streamIdx + sendStreamNum_]));
401 :
402 : // Toggle phase for next slice of same rank
403 2 : recvCurPhase_[streamIdx] = 1 - curPhase;
404 2 : recvDataSlicesByRecvStream_[streamIdx].pop_front();
405 2 : if (recvDataSlicesByRecvStream_[streamIdx].empty()) {
406 2 : nonEmptyRecvStream--;
407 2 : HCCL_INFO("[RunTasks] nonEmptyRecvStream[%u]", nonEmptyRecvStream);
408 : }
409 2 : return HCCL_SUCCESS;
410 2 : }
411 :
412 0 : HcclResult CollBatchSendRecvGroupExecutor::RunTasks(OpParam& param)
413 : {
414 0 : u32 nonEmptySendStream = 0;
415 0 : u32 nonEmptyRecvStream = 0;
416 0 : CHK_RET(CalcStreamTaskStatus(nonEmptySendStream, nonEmptyRecvStream));
417 :
418 0 : CHK_RET(MainPostSubWait(param.stream));
419 :
420 : // Initialize ping-pong state (SDMA only; RDMA不做ping-pong)
421 0 : sendCurPhase_.resize(sendStreamNum_, 0);
422 0 : sendLoadedSize_.resize(sendStreamNum_, 0);
423 0 : sendLoadedRemoteRank_.resize(sendStreamNum_, 0);
424 0 : recvCurPhase_.resize(recvStreamNum_, 0);
425 0 : recvCurRemoteRank_.resize(recvStreamNum_, 0);
426 :
427 0 : u32 pendingSendCount = 0;
428 :
429 0 : while (pendingSendCount > 0 || nonEmptySendStream > 0 || nonEmptyRecvStream > 0 || !rdmaSendSlices_.empty()
430 0 : || !rdmaRecvSlices_.empty()) {
431 0 : HCCL_INFO(
432 : "[RunTasks] pending[%u] sendStream[%u] recvStream[%u] rdmaSend[%zu] rdmaRecv[%zu]", pendingSendCount,
433 : nonEmptySendStream, nonEmptyRecvStream, rdmaSendSlices_.size(), rdmaRecvSlices_.size());
434 :
435 0 : for (u32 i = 0; i < sendStreamNum_; i++) {
436 0 : if (sendLoadedSize_[i] > 0) {
437 0 : CHK_RET(ProcessPreloadedSendSlice(i, pendingSendCount, nonEmptySendStream));
438 0 : } else if (!sendDataSlicesBySendStream_[i].empty()) {
439 0 : CHK_RET(ProcessNewRankSendSlice(i, pendingSendCount, nonEmptySendStream));
440 : }
441 : }
442 :
443 0 : for (u32 i = 0; i < recvStreamNum_; i++) {
444 0 : if (recvDataSlicesByRecvStream_[i].empty()) {
445 0 : continue;
446 : }
447 : // SDMA任务:走ping-pong
448 0 : CHK_RET(ProcessRecvSlice(i, nonEmptyRecvStream));
449 : }
450 :
451 0 : if (!rdmaSendSlices_.empty()) {
452 0 : CHK_RET(ProcessRdmaSendSlice());
453 : }
454 0 : if (!rdmaRecvSlices_.empty()) {
455 0 : CHK_RET(ProcessRdmaRecvSlice());
456 : }
457 :
458 0 : CHK_RET(LaunchTaskExtend(dispatcher_, param.stream, algResResp_->slaveStreams));
459 : }
460 0 : CHK_RET(MainWaitSubPost(param.stream));
461 0 : CHK_RET(LaunchTaskExtend(dispatcher_, param.stream, algResResp_->slaveStreams));
462 0 : return HCCL_SUCCESS;
463 : }
464 :
465 1 : HcclResult CollBatchSendRecvGroupExecutor::ProcessRdmaSendSlice()
466 : {
467 1 : SendRecvSlice& slice = rdmaSendSlices_.front();
468 : // RDMA send使用CCLOut A半区(单流,整半区作为单slot):send scratch offset = 0
469 1 : const u64 sendScratchOffset = 0;
470 1 : DeviceMem sendScratchMem = algResResp_->cclOutputMem.range(sendScratchOffset, slice.size);
471 1 : DeviceMem inMem(slice.addr, slice.size);
472 1 : HCCL_INFO(
473 : "[RunTasks] RdmaSend D2D user[%p] -> CCLOut_A[%p] offset[%llu] size[%llu] rank[%u]", inMem.ptr(),
474 : sendScratchMem.ptr(), sendScratchOffset, slice.size, slice.remoteRank);
475 1 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, sendScratchMem, inMem, algResResp_->slaveStreams[RdmaSendStreamIdx()]));
476 :
477 : // Tx: 本地CCLOut_A(单slot) -> 远端CCLOut_B(单slot)。
478 : // 远端RDMA recv为单流,其recv scratch offset = rdmaDataBlockSize_(B半区基址)。
479 1 : const u64 remoteDstOffset = rdmaDataBlockSize_;
480 1 : LINK sendTargetLink;
481 1 : CHK_RET(GetSendTargetLink(slice.remoteRank, sendTargetLink));
482 1 : CHK_RET(sendTargetLink->TxPrepare(algResResp_->slaveStreams[RdmaSendStreamIdx()]));
483 1 : CHK_RET(sendTargetLink->TxData(
484 : UserMemType::OUTPUT_MEM, remoteDstOffset, sendScratchMem.ptr(), slice.size,
485 : algResResp_->slaveStreams[RdmaSendStreamIdx()]));
486 1 : CHK_RET(sendTargetLink->TxDone(algResResp_->slaveStreams[RdmaSendStreamIdx()]));
487 :
488 1 : rdmaSendSlices_.pop_front();
489 1 : return HCCL_SUCCESS;
490 1 : }
491 :
492 0 : HcclResult CollBatchSendRecvGroupExecutor::ProcessRdmaRecvSlice()
493 : {
494 0 : SendRecvSlice& slice = rdmaRecvSlices_.front();
495 : // RDMA recv使用CCLOut B半区(单流,整半区作为单slot):recv scratch offset = rdmaDataBlockSize_(B半区基址)。
496 : // 与发送方写入的远端offset一致。
497 0 : const u64 recvScratchOffset = rdmaDataBlockSize_;
498 0 : DeviceMem recvScratchMem = algResResp_->cclOutputMem.range(recvScratchOffset, slice.size);
499 :
500 0 : LINK recvTargetLink;
501 0 : CHK_RET(GetRecvTargetLink(slice.remoteRank, recvTargetLink));
502 0 : CHK_RET(recvTargetLink->RxPrepare(algResResp_->slaveStreams[RdmaRecvStreamIdx()]));
503 0 : CHK_RET(recvTargetLink->RxData(
504 : UserMemType::OUTPUT_MEM, recvScratchOffset, recvScratchMem.ptr(), slice.size,
505 : algResResp_->slaveStreams[RdmaRecvStreamIdx()]));
506 0 : CHK_RET(recvTargetLink->RxDone(algResResp_->slaveStreams[RdmaRecvStreamIdx()]));
507 :
508 : // D2D: CCLOut_B -> user output
509 0 : DeviceMem outMem(slice.addr, slice.size);
510 0 : HCCL_INFO(
511 : "[RunTasks] RdmaRecv D2D CCLOut_B[%p] offset[%llu] -> user[%p] size[%llu] rank[%u]", recvScratchMem.ptr(),
512 : recvScratchOffset, outMem.ptr(), slice.size, slice.remoteRank);
513 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, outMem, recvScratchMem, algResResp_->slaveStreams[RdmaRecvStreamIdx()]));
514 :
515 0 : rdmaRecvSlices_.pop_front();
516 0 : return HCCL_SUCCESS;
517 0 : }
518 :
519 1 : HcclResult CollBatchSendRecvGroupExecutor::SetNormalModeIfDeviceDirect()
520 : {
521 3 : for (const auto& q : sendDataSlicesBySendStream_) {
522 3 : for (const auto& slice : q) {
523 1 : LINK targetLink;
524 1 : CHK_RET(GetSendTargetLink(slice.remoteRank, targetLink));
525 1 : if (targetLink->GetTransportType() == TransportType::TRANS_TYPE_DEVICE_DIRECT) {
526 0 : CHK_RET(SetNormalMode(dispatcher_));
527 0 : HCCL_INFO("[CollBatchSendRecvGroupExecutor]Send Set dispatcher NormalMode");
528 0 : return HCCL_SUCCESS;
529 : }
530 1 : }
531 : }
532 :
533 1 : for (const auto& q : recvDataSlicesByRecvStream_) {
534 0 : for (const auto& slice : q) {
535 0 : LINK targetLink;
536 0 : CHK_RET(GetRecvTargetLink(slice.remoteRank, targetLink));
537 0 : if (targetLink->GetTransportType() == TransportType::TRANS_TYPE_DEVICE_DIRECT) {
538 0 : CHK_RET(SetNormalMode(dispatcher_));
539 0 : HCCL_INFO("[CollBatchSendRecvGroupExecutor]Recv Set NormalMode dispatcher");
540 0 : return HCCL_SUCCESS;
541 : }
542 0 : }
543 : }
544 :
545 1 : for (const auto& slice : rdmaSendSlices_) {
546 0 : LINK targetLink;
547 0 : CHK_RET(GetSendTargetLink(slice.remoteRank, targetLink));
548 0 : if (targetLink->GetTransportType() == TransportType::TRANS_TYPE_DEVICE_DIRECT) {
549 0 : CHK_RET(SetNormalMode(dispatcher_));
550 0 : HCCL_INFO("[CollBatchSendRecvGroupExecutor]RdmaSend Set NormalMode dispatcher");
551 0 : return HCCL_SUCCESS;
552 : }
553 0 : }
554 :
555 1 : for (const auto& slice : rdmaRecvSlices_) {
556 0 : LINK targetLink;
557 0 : CHK_RET(GetRecvTargetLink(slice.remoteRank, targetLink));
558 0 : if (targetLink->GetTransportType() == TransportType::TRANS_TYPE_DEVICE_DIRECT) {
559 0 : CHK_RET(SetNormalMode(dispatcher_));
560 0 : HCCL_INFO("[CollBatchSendRecvGroupExecutor]RdmaRecv Set NormalMode dispatcher");
561 0 : return HCCL_SUCCESS;
562 : }
563 0 : }
564 1 : return HCCL_SUCCESS;
565 : }
566 :
567 0 : HcclResult CollBatchSendRecvGroupExecutor::RunLoop(OpParam& param)
568 : {
569 0 : if (static_cast<bool>(topoMatcher_->GetExternalInputHcclEnableFfts())) {
570 0 : auto meta = HcclOpMetaInfo::GetOneForBatchSendRecv();
571 0 : CHK_RET(InitTask(dispatcher_, param.stream, meta.isEnableCache, meta.GetCacheKey()));
572 : // 多流子图前后需加空拷贝
573 0 : CHK_RET(AlgTemplateBase::ExecEmptyTask(
574 : algResResp_->cclInputMem, algResResp_->cclOutputMem, param.stream, dispatcher_));
575 : }
576 0 : CHK_RET(SetNormalModeIfDeviceDirect());
577 0 : CHK_RET(RunTasks(param));
578 0 : if (static_cast<bool>(topoMatcher_->GetExternalInputHcclEnableFfts())) {
579 : // 多流子图前后需加空拷贝
580 0 : CHK_RET(AlgTemplateBase::ExecEmptyTask(
581 : algResResp_->cclInputMem, algResResp_->cclOutputMem, param.stream, dispatcher_));
582 0 : CHK_RET(LaunchTaskExtend(dispatcher_, param.stream, algResResp_->slaveStreams));
583 0 : HCCL_INFO("LaunchTaskExtend!");
584 : }
585 0 : return HCCL_SUCCESS;
586 : }
587 :
588 3 : HcclResult CollBatchSendRecvGroupExecutor::CalcSendSlices()
589 : {
590 : // SDMA slice按remoteRank % streamNum分发到sendDataSlicesBySendStream_(CCLIn, ping-pong);
591 : // RDMA slice按rank分组到rdmaByRank,随后按对称场景规则排序到rdmaSendSlices_(CCLOut A半区, 无ping-pong)。
592 3 : sendDataSlicesBySendStream_.resize(sendStreamNum_);
593 3 : std::map<u32, std::deque<SendRecvSlice>> rdmaByRank;
594 15 : for (u32 i = 0; i < sendStreamNum_; i++) {
595 12 : const auto& sendQueueInner = sendQueueBySendstream_[i];
596 14 : for (u32 j = 0; j < sendQueueInner.size(); j++) {
597 2 : HcclSendRecvItem* sendRecvItem = sendQueueInner[j];
598 2 : u32 unitSize = SIZE_TABLE[sendRecvItem->dataType];
599 2 : bool isRdma = IsRemoteRankRdma(sendRecvItem->remoteRank);
600 2 : u64 maxCountPerLoop = isRdma ? (rdmaDataBlockSize_ / unitSize) : CalcSendLoopMaxCount(unitSize);
601 4 : while (sendRecvItem->count > 0) {
602 2 : u8* curInputPtr = static_cast<u8*>(sendRecvItem->buf);
603 2 : CHK_PTR_NULL(curInputPtr);
604 2 : u64 curCount = (sendRecvItem->count > maxCountPerLoop) ? maxCountPerLoop : sendRecvItem->count;
605 2 : u64 curSize = curCount * unitSize;
606 2 : SendRecvSlice slice(curInputPtr, curSize, sendRecvItem->remoteRank, isRdma);
607 2 : if (isRdma) {
608 1 : rdmaByRank[sendRecvItem->remoteRank].push_back(slice);
609 : } else {
610 1 : sendDataSlicesBySendStream_[i].push_back(slice);
611 : }
612 2 : sendRecvItem->count -= curCount;
613 2 : sendRecvItem->buf = static_cast<u8*>(sendRecvItem->buf) + curSize;
614 : }
615 : }
616 : }
617 : // RDMA按对称场景规则排序(send前向递增,跳过本pod与无任务对端)
618 3 : OrderRdmaSlices(true, rdmaByRank, rdmaSendSlices_);
619 :
620 15 : for (u32 i = 0; i < sendStreamNum_; i++) {
621 12 : const auto& sendQueueInner = sendDataSlicesBySendStream_[i];
622 13 : for (const auto& slice : sendQueueInner) {
623 1 : HCCL_INFO(
624 : "[CalcSendSlices] sendstream[%u] addr[%p] size[%llu] rank[%u] isRdma[%d]", i, slice.addr, slice.size,
625 : slice.remoteRank, slice.isRdma);
626 : }
627 : }
628 4 : for (const auto& slice : rdmaSendSlices_) {
629 1 : HCCL_INFO("[CalcSendSlices] rdmaSend addr[%p] size[%llu] rank[%u]", slice.addr, slice.size, slice.remoteRank);
630 : }
631 3 : return HCCL_SUCCESS;
632 3 : }
633 :
634 3 : HcclResult CollBatchSendRecvGroupExecutor::CalcRecvSlices()
635 : {
636 : // SDMA slice按remoteRank % streamNum分发到recvDataSlicesByRecvStream_(CCLIn, ping-pong);
637 : // RDMA slice按rank分组到rdmaByRank,随后按对称场景规则排序到rdmaRecvSlices_(CCLOut B半区, 无ping-pong)。
638 3 : recvDataSlicesByRecvStream_.resize(recvStreamNum_);
639 3 : std::map<u32, std::deque<SendRecvSlice>> rdmaByRank;
640 15 : for (u32 i = 0; i < recvStreamNum_; i++) {
641 12 : const auto& recvQueueInner = recvQueueByRecvstream_[i];
642 14 : for (u32 j = 0; j < recvQueueInner.size(); j++) {
643 2 : HcclSendRecvItem* sendRecvItem = recvQueueInner[j];
644 2 : u32 unitSize = SIZE_TABLE[sendRecvItem->dataType];
645 2 : bool isRdma = IsRemoteRankRdma(sendRecvItem->remoteRank);
646 2 : u64 maxCountPerLoop = isRdma ? (rdmaDataBlockSize_ / unitSize) : CalcRecvLoopMaxCount(unitSize);
647 4 : while (sendRecvItem->count > 0) {
648 2 : u8* curOutputPtr = static_cast<u8*>(sendRecvItem->buf);
649 2 : CHK_PTR_NULL(curOutputPtr);
650 2 : u64 curCount = (sendRecvItem->count > maxCountPerLoop) ? maxCountPerLoop : sendRecvItem->count;
651 2 : u64 curSize = curCount * unitSize;
652 2 : SendRecvSlice slice(curOutputPtr, curSize, sendRecvItem->remoteRank, isRdma);
653 2 : if (isRdma) {
654 1 : rdmaByRank[sendRecvItem->remoteRank].push_back(slice);
655 : } else {
656 1 : recvDataSlicesByRecvStream_[i].push_back(slice);
657 : }
658 2 : sendRecvItem->count -= curCount;
659 2 : sendRecvItem->buf = static_cast<u8*>(sendRecvItem->buf) + curSize;
660 : }
661 : }
662 : }
663 : // RDMA按对称场景规则排序(recv后向递减,跳过本pod与无任务对端)
664 3 : OrderRdmaSlices(false, rdmaByRank, rdmaRecvSlices_);
665 :
666 15 : for (u32 i = 0; i < recvStreamNum_; i++) {
667 12 : const auto& recvQueueInner = recvDataSlicesByRecvStream_[i];
668 13 : for (const auto& slice : recvQueueInner) {
669 1 : HCCL_INFO(
670 : "[CalcRecvSlices] recvstream[%u] addr[%p] size[%llu] rank[%u] isRdma[%d]", i, slice.addr, slice.size,
671 : slice.remoteRank, slice.isRdma);
672 : }
673 : }
674 4 : for (const auto& slice : rdmaRecvSlices_) {
675 1 : HCCL_INFO("[CalcRecvSlices] rdmaRecv addr[%p] size[%llu] rank[%u]", slice.addr, slice.size, slice.remoteRank);
676 : }
677 3 : return HCCL_SUCCESS;
678 3 : }
679 :
680 64 : u32 CollBatchSendRecvGroupExecutor::GetNextDstRank(u32& curDstRank)
681 : {
682 : // 对称场景send方向:沿rank id递增环绕遍历,跳过本pod。移植自alltoallv_direct_fullmesh。
683 64 : if (curDstRank >= topoAttr_.userRankSize) {
684 5 : curDstRank = curDstRank % topoAttr_.userRankSize;
685 : }
686 64 : if (curDstRank == podStartRank_) {
687 8 : curDstRank += devNumInlocalPod_;
688 : }
689 64 : curDstRank = curDstRank % topoAttr_.userRankSize;
690 64 : return curDstRank++;
691 : }
692 :
693 42 : u32 CollBatchSendRecvGroupExecutor::GetPreSrcRank(u32& curSrcRank)
694 : {
695 : // 对称场景recv方向:沿rank id递减环绕遍历,跳过本pod。移植自alltoallv_direct_fullmesh。
696 42 : if (curSrcRank == podStartRank_ + devNumInlocalPod_ - 1) {
697 2 : curSrcRank = (curSrcRank + topoAttr_.userRankSize - devNumInlocalPod_) % topoAttr_.userRankSize;
698 : }
699 42 : if (curSrcRank == 0) {
700 5 : curSrcRank = topoAttr_.userRankSize - 1;
701 5 : return 0;
702 : }
703 37 : return curSrcRank--;
704 : }
705 :
706 11 : void CollBatchSendRecvGroupExecutor::OrderRdmaSlices(
707 : bool isSend, const std::map<u32, std::deque<SendRecvSlice>>& byRank, std::deque<SendRecvSlice>& out)
708 : {
709 : // 跨pod对端总数 = userRankSize - devNumInlocalPod。遍历每个候选rank一次。
710 11 : u32 totalRdmaRankNum = topoAttr_.userRankSize - devNumInlocalPod_;
711 11 : u32 curRank = INVALID_VALUE_RANKID;
712 11 : if (isSuPodAsym_) {
713 : // 非对称场景:send/recv使用相同遍历顺序(均前向递增GetNextDstRank)。
714 : // 起点为本pod之外的第一个rank(从0开始扫描,取首个不在[podStartRank_, podEndRank_]内的rank)。
715 0 : for (u32 i = 0; i < topoAttr_.userRankSize; i++) {
716 0 : if (i < podStartRank_ || i > podEndRank_) {
717 0 : curRank = i;
718 0 : break;
719 : }
720 : }
721 0 : HCCL_INFO(
722 : "[OrderRdmaSlices] %s asym startRank[%u] totalRdmaRankNum[%u]", isSend ? "send" : "recv", curRank,
723 : totalRdmaRankNum);
724 0 : for (u32 i = 0; i < totalRdmaRankNum; i++) {
725 0 : u32 rank = GetNextDstRank(curRank);
726 0 : auto it = byRank.find(rank);
727 0 : if (it == byRank.end()) {
728 0 : HCCL_INFO("[OrderRdmaSlices] %s asym skip rank[%u] (no task)", isSend ? "send" : "recv", rank);
729 0 : continue;
730 : }
731 0 : for (const auto& s : it->second) {
732 0 : out.push_back(s);
733 : }
734 0 : HCCL_INFO(
735 : "[OrderRdmaSlices] %s asym append rank[%u] sliceNum[%zu]", isSend ? "send" : "recv", rank,
736 : it->second.size());
737 : }
738 : } else {
739 : // 对称场景:send前向递增(GetNextDstRank),recv后向递减(GetPreSrcRank)。
740 : // send起点取"下一个pod中相同pod内位置的rank";recv起点取"上一个pod中相同pod内位置的rank"。
741 11 : curRank = isSend ? (topoAttr_.userRank + devNumInlocalPod_) % topoAttr_.userRankSize :
742 4 : (topoAttr_.userRank + topoAttr_.userRankSize - devNumInlocalPod_) % topoAttr_.userRankSize;
743 11 : HCCL_INFO(
744 : "[OrderRdmaSlices] %s sym startRank[%u] totalRdmaRankNum[%u]", isSend ? "send" : "recv", curRank,
745 : totalRdmaRankNum);
746 103 : for (u32 i = 0; i < totalRdmaRankNum; i++) {
747 92 : u32 rank = isSend ? GetNextDstRank(curRank) : GetPreSrcRank(curRank);
748 92 : auto it = byRank.find(rank);
749 92 : if (it == byRank.end()) {
750 80 : HCCL_INFO("[OrderRdmaSlices] %s sym skip rank[%u] (no task)", isSend ? "send" : "recv", rank);
751 80 : continue;
752 : }
753 25 : for (const auto& s : it->second) {
754 13 : out.push_back(s);
755 : }
756 12 : HCCL_INFO(
757 : "[OrderRdmaSlices] %s sym append rank[%u] sliceNum[%zu]", isSend ? "send" : "recv", rank,
758 : it->second.size());
759 : }
760 : }
761 11 : }
762 :
763 3 : u64 CollBatchSendRecvGroupExecutor::CalcSendLoopMaxCount(const u32 unitSize) const
764 : {
765 : // 中转内存单次最多能够接受的input count
766 3 : u64 maxCountPerLoop = bufferSliceSize_ / unitSize;
767 3 : HCCL_INFO(
768 : "[CollBatchSendRecvGroupExecutor][CalcSendLoopMaxCount]"
769 : "using default maxCountPerLoop[%llu] as CCLBuffSize / unitSize.",
770 : maxCountPerLoop);
771 3 : return maxCountPerLoop;
772 : }
773 :
774 3 : u64 CollBatchSendRecvGroupExecutor::CalcRecvLoopMaxCount(const u32 unitSize) const
775 : {
776 : // 中转内存单次最多能够接受的output count
777 3 : u64 maxCountPerLoop = bufferSliceSize_ / unitSize;
778 3 : HCCL_INFO(
779 : "[CollBatchSendRecvGroupExecutor][CalcRecvLoopMaxCount]"
780 : "using default maxCountPerLoop[%llu] as CCLBuffSize / unitSize.",
781 : maxCountPerLoop);
782 3 : return maxCountPerLoop;
783 : }
784 :
785 1 : HcclResult CollBatchSendRecvGroupExecutor::CalcStreamNum(u32& streamNum)
786 : {
787 1 : if (topoAttr_.userRankSize == 1) {
788 0 : sendStreamNum_ = 0;
789 0 : recvStreamNum_ = 0;
790 0 : streamNum = 0;
791 0 : HCCL_INFO("[CollBatchSendRecvGroupExecutor] Only one rank, do not need substream, streamNum[%u]", streamNum);
792 0 : return HCCL_SUCCESS;
793 : }
794 1 : sendStreamNum_ = GROUP_MAX_CONCURRENT;
795 1 : recvStreamNum_ = GROUP_MAX_CONCURRENT;
796 : // SDMA占sendStreamNum_+recvStreamNum_条从流;RDMA单独占RDMA_STREAM_NUM条(1 send + 1 recv)。
797 1 : streamNum = sendStreamNum_ + recvStreamNum_ + RDMA_STREAM_NUM;
798 1 : HCCL_INFO("[CollBatchSendRecvGroupExecutor][CalcStreamNum] tag_[%s], streamNum[%u].", tag_.c_str(), streamNum);
799 1 : return HCCL_SUCCESS;
800 : }
801 :
802 : REGISTER_EXEC("BatchSendRecvGroup", BatchSendRecvGroupExecutor, CollBatchSendRecvGroupExecutor);
803 : } // namespace hccl
|