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 "allltoall_pipeline_mesh_pairwise_ccl_enough.h"
12 : #include "alg_template_register.h"
13 :
14 : namespace hccl {
15 :
16 : static const u32 INTRA_STREAM_INFO_SENDLEN_INDEX = 0; // intraStreamInfo 中 sendLen 的下标
17 : static const u32 INTRA_STREAM_INFO_RECVLEN_INDEX = 1; // intraStreamInfo 中 recvLen 的下标
18 : static const u32 INTRA_STREAM_INFO_RECV_REMOTE_OFFSET_INDEX = 2; // intraStreamInfo 中 recvRemoteOffset 的下标
19 : static const u32 INTRA_STREAM_INFO_RECV_LOCAL_OFFSET_INDEX = 3; // intraStreamInfo 中 recvLocalOffset 的下标
20 :
21 0 : AlltoallPipelineMeshPairwiseCCLEnough::AlltoallPipelineMeshPairwiseCCLEnough(
22 0 : const HcclDispatcher dispatcher): AlltoallPipelineBase(dispatcher) {}
23 :
24 0 : AlltoallPipelineMeshPairwiseCCLEnough::~AlltoallPipelineMeshPairwiseCCLEnough() {}
25 :
26 0 : u32 AlltoallPipelineMeshPairwiseCCLEnough::CalcInterNumSteps()
27 : {
28 0 : return interRankSize_ - 1;
29 : }
30 :
31 : // 适配新CollExecutor接口
32 0 : HcclResult AlltoallPipelineMeshPairwiseCCLEnough::Prepare(u32 userRank, A2aPipelineMemory A2aPipelineMemory,
33 : const SubCommInfo &level0CommInfo, const SubCommInfo &level1CommInfo,
34 : Stream &mainStream, std::vector<Stream> &subStream,
35 : std::vector<std::shared_ptr<LocalNotify>> ¬ifyMain, std::vector<std::shared_ptr<LocalNotify>> ¬ifySub,
36 : std::vector<SendRecvInfo> &allMeshAggregationSendRecvInfo, HcclWorkflowMode workMode)
37 : {
38 0 : AlltoallPipelineBase::Prepare(userRank, A2aPipelineMemory, level0CommInfo, level1CommInfo, mainStream,
39 : subStream, notifyMain, notifySub, allMeshAggregationSendRecvInfo, workMode);
40 0 : GetIntraScratchOffset();
41 0 : CHK_RET(DeviceMemMapping());
42 0 : return HCCL_SUCCESS;
43 : }
44 :
45 : // 统一计算每步 mesh 内收发时从各卡 scratch 读取的 offset 和 length
46 0 : HcclResult AlltoallPipelineMeshPairwiseCCLEnough::GetIntraScratchOffset()
47 : {
48 0 : for (u32 i = 0; i < intraRankSize_; i++) {
49 0 : intraScratchOffsetMap_[i] = std::vector<u64>();
50 0 : intraScratchLengMap_[i] = std::vector<u64>();
51 0 : u64 startOffset = 0;
52 0 : for (u32 remoteRank = i; remoteRank < groupRankSize_; remoteRank += intraRankSize_) {
53 0 : if (remoteRank == userRank_) {
54 0 : localScratchOffset_ = startOffset;
55 : }
56 0 : const std::vector<u64>& remoteSendOffset = (*allMeshAggregationSendRecvInfo_)[remoteRank].sendOffset;
57 0 : const std::vector<u64>& remoteSendLength = (*allMeshAggregationSendRecvInfo_)[remoteRank].sendLength;
58 0 : intraScratchOffsetMap_[i].push_back(startOffset + (remoteSendOffset[userRank_] -
59 0 : remoteSendOffset[meshRankStart_]));
60 0 : startOffset += (remoteSendOffset[meshRankEnd_] + remoteSendLength[meshRankEnd_] -
61 0 : remoteSendOffset[meshRankStart_]);
62 0 : intraScratchLengMap_[i].push_back(remoteSendLength[userRank_]);
63 : }
64 : }
65 0 : return HCCL_SUCCESS;
66 : }
67 :
68 0 : HcclResult AlltoallPipelineMeshPairwiseCCLEnough::DeviceMemMapping()
69 : {
70 0 : if (workMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
71 0 : interTransportSend_ = cclIn_;
72 0 : interTransportRecv_ = cclOut_;
73 0 : intraTransportSend_ = cclOut_;
74 : } else {
75 0 : interTransportSend_ = inputMem_;
76 0 : interTransportRecv_ = scratchMem_;
77 0 : intraTransportSend_ = scratchMem_;
78 : }
79 0 : for (u32 intraRank = 0; intraRank < intraRankSize_; intraRank++) {
80 0 : if (intraRank == intraRankId_) {
81 0 : continue;
82 : }
83 0 : LINK& intraNeighboorTransport = intraLinks_[intraRank];
84 0 : void* remDMAMemPtr = nullptr;
85 0 : CHK_RET(intraNeighboorTransport->GetRemoteMem(UserMemType::INPUT_MEM, &remDMAMemPtr));
86 : DeviceMem remoteAlltoallScratch = DeviceMem::create(static_cast<u8 *>(remDMAMemPtr),
87 0 : workMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE ? cclIn_.size() : scratchMem_.size());
88 0 : intraNeighBoorMemory_[intraRank] = {remoteAlltoallScratch};
89 0 : }
90 0 : return HCCL_SUCCESS;
91 0 : }
92 :
93 0 : HcclResult AlltoallPipelineMeshPairwiseCCLEnough::PrepareInterData(u32 step)
94 : {
95 : // 准备 mesh 间发送信息
96 0 : nextInterSendData_.clear();
97 0 : u32 interSendRankStart = ((interRankId_ + 1 + step) % interRankSize_) * intraRankSize_;
98 0 : u32 interSendRankEnd = interSendRankStart + intraRankSize_ - 1;
99 0 : u64 startMemOffset = localSendRecvInfo_.sendOffset[interSendRankStart];
100 0 : u64 meshSendLength = localSendRecvInfo_.sendOffset[interSendRankEnd] +
101 0 : localSendRecvInfo_.sendLength[interSendRankEnd] - startMemOffset;
102 0 : u64 sendDestOffset = 0;
103 0 : for (u32 relatedRank = intraRankId_; relatedRank < userRank_; relatedRank += intraRankSize_) {
104 0 : const SendRecvInfo& info = (*allMeshAggregationSendRecvInfo_)[relatedRank];
105 0 : sendDestOffset += (info.sendOffset[interSendRankEnd] + info.sendLength[interSendRankEnd] -
106 0 : info.sendOffset[interSendRankStart]);
107 : }
108 0 : DeviceMem srcMem = inputMem_.range(startMemOffset, meshSendLength);
109 0 : DeviceMem dstMem = interTransportSend_.range(startMemOffset, meshSendLength);
110 0 : HCCL_DEBUG("[AlltoallPipelineMeshPairwiseCCLEnough][PrepareInterSendData] userRank %u, interRank %u, "
111 : "intraRank %u move from userInput offset %llu length %llu to interTransportSend, send to remote offset "
112 : "%llu", userRank_, interRankId_, intraRankId_, startMemOffset, meshSendLength, sendDestOffset);
113 :
114 0 : HCCL_DEBUG("user size %u, inter size %u, intra size %u",
115 : groupRankSize_, interRankSize_, intraRankSize_);
116 0 : if (workMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
117 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dstMem, srcMem, mainStream_));
118 : }
119 0 : nextInterSendData_.emplace_back(TxMemoryInfo{UserMemType::OUTPUT_MEM, sendDestOffset, dstMem.ptr(),
120 : meshSendLength});
121 :
122 : // 准备 mesh 间接收信息
123 0 : nextInterRecvData_.clear();
124 0 : u32 recvBlockStart = (((interRankId_ + interRankSize_ - 1u - step) % interRankSize_) * intraRankSize_);
125 0 : u64 recvLocalOffset = 0;
126 0 : for (u32 relatedRank = intraRankId_; relatedRank < recvBlockStart; relatedRank += intraRankSize_) {
127 0 : const SendRecvInfo& info = (*allMeshAggregationSendRecvInfo_)[relatedRank];
128 0 : recvLocalOffset += (info.sendOffset[meshRankEnd_] + info.sendLength[meshRankEnd_] -
129 0 : info.sendOffset[meshRankStart_]);
130 : }
131 0 : const SendRecvInfo& recvRankInfo = (*allMeshAggregationSendRecvInfo_)[recvBlockStart + intraRankId_];
132 0 : u64 recvRemoteOffset = recvRankInfo.sendOffset[meshRankStart_];
133 0 : u64 recvLength = (recvRankInfo.sendOffset[meshRankEnd_] + recvRankInfo.sendLength[meshRankEnd_] -
134 0 : recvRankInfo.sendOffset[meshRankStart_]);
135 0 : nextInterRecvData_.emplace_back(RxMemoryInfo{UserMemType::INPUT_MEM, recvRemoteOffset,
136 0 : interTransportRecv_.range(recvLocalOffset, recvLength).ptr(), recvLength});
137 0 : return HCCL_SUCCESS;
138 0 : }
139 :
140 : // 将原先在userInput,且需要发到本mesh内其他卡的数据搬到CCL
141 0 : HcclResult AlltoallPipelineMeshPairwiseCCLEnough::PrepareIntraData()
142 : {
143 0 : u64 startMemOffset = localSendRecvInfo_.sendOffset[meshRankStart_];
144 0 : u64 meshSendLength = localSendRecvInfo_.sendOffset[meshRankEnd_] +
145 0 : localSendRecvInfo_.sendLength[meshRankEnd_] - startMemOffset;
146 0 : u64 intraSendOffset = 0;
147 0 : for (u32 relatedRank = intraRankId_; relatedRank < meshRankStart_; relatedRank += intraRankSize_) {
148 0 : const SendRecvInfo& info = (*allMeshAggregationSendRecvInfo_)[relatedRank];
149 0 : intraSendOffset += (info.sendOffset[meshRankEnd_] + info.sendLength[meshRankEnd_] -
150 0 : info.sendOffset[meshRankStart_]);
151 : }
152 0 : DeviceMem srcMem = inputMem_.range(startMemOffset, meshSendLength);
153 0 : DeviceMem dstMem = intraTransportSend_.range(intraSendOffset, meshSendLength);
154 0 : if (workMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
155 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dstMem, srcMem, mainStream_));
156 : }
157 0 : HCCL_DEBUG("[AlltoallPipelineMeshPairwiseCCLEnough][PrepareIntraData] userRank %u, interRank %u, intraRank %u "
158 : "copy from userInput offset %llu length %llu to intraTransportSend offset %llu", userRank_, interRankId_,
159 : intraRankId_, startMemOffset, meshSendLength, intraSendOffset);
160 0 : return HCCL_SUCCESS;
161 0 : }
162 :
163 0 : void AlltoallPipelineMeshPairwiseCCLEnough::UpdateIntraStreamInfo(u32 step)
164 : {
165 0 : intraStreamInfo_.clear();
166 0 : u32 localMeshIndex = (interRankId_ + interRankSize_ - step) % interRankSize_;
167 0 : u32 firstDataBlockIndex = (meshRankStart_ + groupRankSize_ - step * intraRankSize_) % groupRankSize_;
168 0 : const std::vector<u64>& sendLengths = (*allMeshAggregationSendRecvInfo_)[firstDataBlockIndex +
169 0 : intraRankId_].sendLength;
170 0 : const std::vector<u64>& recvLengths = localSendRecvInfo_.recvLength;
171 0 : const std::vector<u64>& recvOffsets = localSendRecvInfo_.recvOffset;
172 0 : HCCL_DEBUG("[AlltoallPipelineMeshPairwiseCCLEnough][UpdateIntraStreamInfo] userRank %u, "
173 : "interRank %u, intraRank %u, step %u", userRank_, interRankId_, intraRankId_, step);
174 0 : for (u32 intraRank = 0; intraRank < intraRankSize_; intraRank++) {
175 0 : u64 sendLen = sendLengths[meshRankStart_ + intraRank];
176 0 : u64 recvLen = recvLengths[firstDataBlockIndex + intraRank];
177 0 : u64 recvRemoteOffset = intraScratchOffsetMap_[intraRank][localMeshIndex];
178 0 : u64 recvLocalOffset = recvOffsets[firstDataBlockIndex + intraRank];
179 0 : if (intraRank != intraRankId_) {
180 0 : intraStreamInfo_[intraRank] = {sendLen, recvLen, recvRemoteOffset, recvLocalOffset};
181 0 : HCCL_DEBUG("[AlltoallPipelineMeshPairwiseCCLEnough][UpdateIntraStreamInfo] userRank %u, interRank %u, "
182 : "intraRank %u, sdma stream %u need send %llu and read length %llu from remote offset %llu "
183 : "to local offset %llu", userRank_, interRankId_, intraRankId_, intraRank, sendLen, recvLen,
184 : recvRemoteOffset, recvLocalOffset);
185 : }
186 : }
187 0 : }
188 :
189 0 : HcclResult AlltoallPipelineMeshPairwiseCCLEnough::SendRecvDataIntraMesh()
190 : {
191 0 : HCCL_DEBUG("[AlltoallPipelineMeshPairwiseCCLEnough][SendRecvDataIntraMesh] userRank %u, "
192 : "interRank %u, intraRank %u, sdma stream %s wait main stream", userRank_, interRankId_,
193 : intraRankId_, GetStreamIndexString().c_str());
194 0 : for (auto& intraInfo : intraStreamInfo_) {
195 0 : u32 streamIndex = intraInfo.first;
196 0 : u64 recvLen = intraInfo.second[INTRA_STREAM_INFO_RECVLEN_INDEX];
197 0 : Stream& currStream = subStream_[streamIndex];
198 0 : LINK& readTransport = intraLinks_[streamIndex];
199 0 : CHK_RET(readTransport->TxAck(currStream));
200 0 : CHK_RET(readTransport->RxAck(currStream));
201 0 : u64 recvRemoteOffset = intraInfo.second[INTRA_STREAM_INFO_RECV_REMOTE_OFFSET_INDEX];
202 0 : u64 recvLocalOffset = intraInfo.second[INTRA_STREAM_INFO_RECV_LOCAL_OFFSET_INDEX];
203 0 : DeviceMem src = intraNeighBoorMemory_[streamIndex][0].range(recvRemoteOffset, recvLen);
204 0 : DeviceMem dst = outputMem_.range(recvLocalOffset, recvLen);
205 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, currStream, readTransport->GetRemoteRank(),
206 : readTransport->GetLinkType()));
207 0 : CHK_RET(readTransport->TxDataSignal(currStream));
208 0 : HCCL_DEBUG("[AlltoallPipelineMeshPairwiseCCLEnough][SendRecvDataIntraMesh] userRank %u, interRank %u, "
209 : "intraRank %u, sdma stream %llu read data from remote offset %llu len %llu to local %llu",
210 : userRank_, interRankId_, intraRankId_, streamIndex, recvRemoteOffset, recvLen, recvLocalOffset);
211 0 : CHK_RET(readTransport->RxDataSignal(currStream));
212 0 : }
213 0 : HCCL_DEBUG("[AlltoallPipelineMeshPairwiseCCLEnough][SendRecvDataIntraMesh] userRank %u, interRank %u, "
214 : "intraRank %u, sdma stream %s notify main stream", userRank_, interRankId_, intraRankId_,
215 : GetStreamIndexString().c_str());
216 0 : return HCCL_SUCCESS;
217 : }
218 :
219 0 : HcclResult AlltoallPipelineMeshPairwiseCCLEnough::SendRecvDataInterMesh(u32 step)
220 : {
221 0 : Stream& interStream = subStream_[intraRankId_];
222 0 : LINK& interRecvTransport = interLinks_[(interRankId_ + interRankSize_ - 1 - step) % interRankSize_];
223 0 : LINK& interSendTransport = interLinks_[(interRankId_ + 1 + step) % interRankSize_];
224 0 : CHK_RET(interRecvTransport->TxAck(interStream));
225 0 : CHK_RET(interSendTransport->RxAck(interStream));
226 0 : CHK_RET(interSendTransport->TxAsync(nextInterSendData_, interStream));
227 0 : CHK_RET(interRecvTransport->RxAsync(nextInterRecvData_, interStream));
228 0 : CHK_RET(interRecvTransport->PostFinAck(interStream));
229 0 : CHK_RET(interSendTransport->WaitFinAck(interStream));
230 0 : CHK_RET(ExecuteBarrier(interRecvTransport, interSendTransport, interStream));
231 0 : return HCCL_SUCCESS;
232 : }
233 :
234 0 : HcclResult AlltoallPipelineMeshPairwiseCCLEnough::LocalCopyDataRecvFromInter(u32 interRankDistance)
235 : {
236 0 : u64 scratchOffset = intraScratchOffsetMap_[intraRankId_][(interRankId_ +
237 0 : interRankSize_ - interRankDistance) % interRankSize_];
238 0 : u64 recvLen = localSendRecvInfo_.recvLength[(userRank_ + groupRankSize_ -
239 0 : interRankDistance * intraRankSize_) % groupRankSize_];
240 0 : u64 userOutOffset = localSendRecvInfo_.recvOffset[(userRank_ + groupRankSize_ -
241 0 : interRankDistance * intraRankSize_) % groupRankSize_];
242 0 : DeviceMem src = interTransportRecv_.range(scratchOffset, recvLen);
243 0 : DeviceMem dst = outputMem_.range(userOutOffset, recvLen);
244 0 : HCCL_DEBUG("[AlltoallPipelineMeshPairwiseCCLEnough][LocalCopyDataRecvFromInter]local move from "
245 : "interTransportRecv_ offset %llu length %llu to outputMem_ %llu", scratchOffset, recvLen, userOutOffset);
246 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, mainStream_));
247 0 : return HCCL_SUCCESS;
248 0 : }
249 :
250 0 : HcclResult AlltoallPipelineMeshPairwiseCCLEnough::PreProcess()
251 : {
252 : // server 间收发时间较长,先搬 server 间收发所需数据然后马上让server间开始收发
253 0 : CHK_RET(PrepareInterData(0));
254 0 : CHK_RET(NotifyInterStreamStart());
255 : // 之后将 server 内所需数据准备好之后唤醒server内从流收发
256 0 : UpdateIntraStreamInfo(0);
257 0 : CHK_RET(PrepareIntraData());
258 0 : CHK_RET(NotifyIntraStreamStart());
259 0 : CHK_RET(SendRecvDataIntraMesh());
260 0 : return HCCL_SUCCESS;
261 : }
262 :
263 0 : HcclResult AlltoallPipelineMeshPairwiseCCLEnough::PipelineSend(u32 step, bool isLastStep)
264 : {
265 0 : CHK_RET(SendRecvDataInterMesh(step));
266 0 : CHK_RET(PrepareInterData(step + 1u));
267 0 : CHK_RET(WaitInterStreamFinish());
268 0 : CHK_RET(WaitIntraStreamFinish());
269 0 : CHK_RET(ExecEmptyTask(inputMem_, outputMem_, mainStream_, dispatcher_));
270 0 : if (!isLastStep) {
271 0 : CHK_RET(NotifyInterStreamStart());
272 : }
273 0 : UpdateIntraStreamInfo(step + 1u);
274 0 : CHK_RET(NotifyIntraStreamStart());
275 0 : CHK_RET(SendRecvDataIntraMesh());
276 0 : CHK_RET(LocalCopyDataRecvFromInter(step + 1u));
277 0 : return HCCL_SUCCESS;
278 : }
279 :
280 0 : HcclResult AlltoallPipelineMeshPairwiseCCLEnough::PostProcess()
281 : {
282 : // 最后的收尾工作
283 0 : CHK_RET(LocalCopyDataRecvFromInter(0));
284 0 : CHK_RET(WaitIntraStreamFinish());
285 0 : return HCCL_SUCCESS;
286 : }
287 : REGISTER_TEMPLATE(TemplateType::TEMPLATE_ALL_2_ALL_PIPELINE_MESH_PAIRWISE_CCL_ENOUGH,
288 : AlltoallPipelineMeshPairwiseCCLEnough);
289 : } // namespace hccl
|