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