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 "alltoallv_for_310p.h"
12 : #include "alg_template_register.h"
13 :
14 : namespace hccl {
15 1 : AlltoAllVFor310P::AlltoAllVFor310P(const HcclDispatcher dispatcher) : AlgTemplateBase(dispatcher) {}
16 :
17 2 : AlltoAllVFor310P::~AlltoAllVFor310P() {}
18 :
19 1 : HcclResult AlltoAllVFor310P::Prepare(
20 : DeviceMem& userInput, DeviceMem& userOutput, DeviceMem& cclInMem, DeviceMem& cclOutMem,
21 : const std::vector<std::shared_ptr<LocalNotify>>& signalMainToSub,
22 : const std::vector<std::shared_ptr<LocalNotify>>& signalSubToMain, Stream& mainStream,
23 : std::vector<Stream>& subStreams, const std::vector<LINK>& links, u32 userRank, u32 userRankSize,
24 : std::vector<SendRecvInfo>& allMeshAggregationSendRecvInfo)
25 : {
26 1 : mainStream_ = mainStream;
27 1 : subStream_ = subStreams;
28 1 : links_ = links;
29 1 : userRank_ = userRank;
30 1 : userRankSize_ = userRankSize;
31 1 : CHK_PRT_RET(userRankSize_ == 0, HCCL_ERROR("[AlltoAllVFor310P][Prepare]userRankSize_ is zero."), HCCL_E_PARA);
32 1 : allMeshAggregationSendRecvInfoPtr_ = &allMeshAggregationSendRecvInfo;
33 :
34 1 : userInput_ = userInput;
35 1 : userOutput_ = userOutput;
36 1 : cclInMem_ = cclInMem;
37 1 : cclOutMem_ = cclOutMem;
38 1 : memList_.push_back(cclInMem_); // Id 0
39 1 : memList_.push_back(cclOutMem_); // Id 1
40 1 : memList_.push_back(userOutput_); // Id 2
41 :
42 1 : if (userRank_ % COMPUTE_CONST == 0) {
43 1 : mainRank_ = true;
44 1 : myMinor_ = userRank_ + 1;
45 1 : if (subStream_.size() != COMPUTE_CONST) {
46 0 : HCCL_ERROR("[AlltoAllVFor310P][Prepare]main subStream.size[%zu] != 2", subStream_.size());
47 0 : return HCCL_E_INTERNAL;
48 : }
49 : } else {
50 0 : minorRank_ = true;
51 0 : myMain_ = (userRank_ - 1 + userRankSize_) % userRankSize_;
52 0 : if (subStream_.size() != 1) {
53 0 : HCCL_ERROR("[AlltoAllVFor310P][Prepare]minor subStream.size[%zu] != 1", subStream_.size());
54 0 : return HCCL_E_INTERNAL;
55 : }
56 : }
57 :
58 1 : CHK_PRT_RET(
59 : signalMainToSub.size() != subStream_.size() || signalSubToMain.size() != subStream_.size(),
60 : HCCL_ERROR(
61 : "[AlltoAllVFor310P][Prepare] Signal size not equal to subStream size, signalMainToSub.size[%llu],"
62 : "signalSubToMain.size[%llu], subStream_.size[%llu]",
63 : signalMainToSub.size(), signalSubToMain.size(), subStream_.size()),
64 : HCCL_E_INTERNAL);
65 :
66 1 : HCCL_DEBUG(
67 : "userRank[%u], subStream.size[%zu], signalMainToSub.size[%zu], signalSubToMain.size[%zu]", userRank_,
68 : subStream_.size(), signalMainToSub.size(), signalSubToMain.size());
69 :
70 3 : for (u32 index = 0; index < signalMainToSub.size(); index++) {
71 2 : CHK_PTR_NULL(signalMainToSub[index]);
72 2 : signalMainToSub_.push_back(signalMainToSub[index]);
73 : }
74 :
75 3 : for (u32 index = 0; index < signalSubToMain.size(); index++) {
76 2 : CHK_PTR_NULL(signalSubToMain[index]);
77 2 : signalSubToMain_.push_back(signalSubToMain[index]);
78 : }
79 :
80 1 : cclBlockSize_ = ((cclInMem.size() / COMPUTE_CONST) / ALIGN_CONST) * ALIGN_CONST; // 除以128取整
81 1 : maxSizePerLoop_ = cclBlockSize_ - ALIGN_CONST;
82 :
83 1 : CHK_PRT_RET(cclBlockSize_ == 0, HCCL_ERROR("[AlltoAllVFor310P][Prepare]DataBlockSize_is zero."), HCCL_E_INTERNAL);
84 :
85 1 : return HCCL_SUCCESS;
86 : }
87 :
88 0 : std::string AlltoAllVFor310P::GetStreamIndexString()
89 : {
90 0 : std::string res = "";
91 0 : for (u32 streamIndex = 0; streamIndex < subStream_.size(); streamIndex++) {
92 0 : res += std::to_string(streamIndex) + ", ";
93 : }
94 0 : return res;
95 0 : }
96 :
97 0 : HcclResult AlltoAllVFor310P::WaitSubStreamFinish()
98 : {
99 : // 从流通知主流做完
100 0 : for (u32 streamIndex = 0; streamIndex < subStream_.size(); streamIndex++) {
101 0 : CHK_RET(LocalNotify::Post(
102 : subStream_[streamIndex], dispatcher_, signalMainToSub_[streamIndex], INVALID_VALUE_STAGE));
103 0 : CHK_RET(LocalNotify::Wait(mainStream_, dispatcher_, signalMainToSub_[streamIndex], INVALID_VALUE_STAGE));
104 : }
105 0 : HCCL_DEBUG(
106 : "[AlltoAllVFor310P][WaitSubStreamFinish] userRank [%u] main stream wait stream [%s]", userRank_,
107 : GetStreamIndexString().c_str());
108 0 : return HCCL_SUCCESS;
109 : }
110 :
111 0 : HcclResult AlltoAllVFor310P::NotifySubStreamStart()
112 : {
113 0 : for (u32 streamIndex = 0; streamIndex < subStream_.size(); streamIndex++) {
114 0 : CHK_RET(LocalNotify::Post(mainStream_, dispatcher_, signalSubToMain_[streamIndex], INVALID_VALUE_STAGE));
115 0 : CHK_RET(LocalNotify::Wait(
116 : subStream_[streamIndex], dispatcher_, signalSubToMain_[streamIndex], INVALID_VALUE_STAGE));
117 : }
118 0 : HCCL_DEBUG(
119 : "[AlltoAllVFor310P][NotifySubStreamStart] userRank [%u] main stream notify sdma stream [%s]", userRank_,
120 : GetStreamIndexString().c_str());
121 0 : return HCCL_SUCCESS;
122 : }
123 :
124 0 : HcclResult AlltoAllVFor310P::CalcSendInfo(
125 : const u32 srcDataRank, const u32 dstDataRank, const u32 times, const u64 subStepLen, SendMemBlock& sendInfo)
126 : {
127 0 : const std::vector<u64>& sendLength = (*allMeshAggregationSendRecvInfoPtr_)[srcDataRank].sendLength;
128 0 : const std::vector<u64>& sendOffset = (*allMeshAggregationSendRecvInfoPtr_)[srcDataRank].sendOffset;
129 0 : const std::vector<u64>& recvOffset = (*allMeshAggregationSendRecvInfoPtr_)[dstDataRank].recvOffset;
130 :
131 0 : u32 sendLen = 0;
132 0 : if (sendLength[dstDataRank] > times * maxSizePerLoop_) {
133 0 : u32 leftLen = sendLength[dstDataRank] - times * maxSizePerLoop_;
134 0 : sendLen = leftLen > subStepLen ? subStepLen : leftLen;
135 0 : sendInfo.userInOffset = sendOffset[dstDataRank] + times * maxSizePerLoop_;
136 : } else {
137 : sendInfo.userInOffset
138 0 : = sendOffset[dstDataRank] + sendLength[dstDataRank]; // 已经发完了,offset变成最大值,sendLen为0
139 : }
140 0 : sendInfo.dstRank = dstDataRank;
141 0 : sendInfo.sendLen = sendLen;
142 0 : sendInfo.cclDstOffset = (recvOffset[srcDataRank] + times * maxSizePerLoop_) % ALIGN_CONST;
143 0 : HCCL_DEBUG(
144 : "[AlltoAllVFor310P] [CalcSendInfo]srcDataRank[%u], dstDataRank[%u], times[%u], subStepLen[%llu], "
145 : "sendLen[%llu], userInOffset[%llu], cclDstOffset[%llu]",
146 : srcDataRank, dstDataRank, times, subStepLen, sendInfo.sendLen, sendInfo.userInOffset, sendInfo.cclDstOffset);
147 :
148 0 : return HCCL_SUCCESS;
149 : }
150 :
151 0 : HcclResult AlltoAllVFor310P::CalcRecvInfo(
152 : const u32 srcDataRank, const u32 dstDataRank, const u32 times, const u64 subStepLen, RecvMemBlock& recvInfo)
153 : {
154 0 : const std::vector<u64>& recvLength = (*allMeshAggregationSendRecvInfoPtr_)[dstDataRank].recvLength;
155 0 : const std::vector<u64>& recvOffset = (*allMeshAggregationSendRecvInfoPtr_)[dstDataRank].recvOffset;
156 :
157 0 : u32 recvLen = 0;
158 0 : if (recvLength[srcDataRank] > times * maxSizePerLoop_) {
159 0 : u32 leftLen = recvLength[srcDataRank] - times * maxSizePerLoop_;
160 0 : recvLen = leftLen > subStepLen ? subStepLen : leftLen;
161 0 : recvInfo.userOutOffset = recvOffset[srcDataRank] + times * maxSizePerLoop_;
162 : } else {
163 : recvInfo.userOutOffset
164 0 : = recvOffset[srcDataRank] + recvLength[srcDataRank]; // 已经收完了,offset变成最大值,recvLen为0
165 : }
166 0 : recvInfo.srcRank = srcDataRank;
167 0 : recvInfo.recvLen = recvLen;
168 0 : recvInfo.cclSrcOffset = (recvOffset[srcDataRank] + times * maxSizePerLoop_) % ALIGN_CONST;
169 0 : HCCL_DEBUG(
170 : "[AlltoAllVFor310P] [CalcRecvInfo]dstDataRank[%u], srcDataRank[%u], times[%u], subStepLen[%llu], "
171 : "recvLen[%llu], userOutOffset[%llu], cclSrcOffset[%llu]",
172 : dstDataRank, srcDataRank, times, subStepLen, recvInfo.recvLen, recvInfo.userOutOffset, recvInfo.cclSrcOffset);
173 :
174 0 : return HCCL_SUCCESS;
175 : }
176 :
177 0 : HcclResult AlltoAllVFor310P::MainFirstLocalCopy(const u32 times, const u32 roundIdx, const u64 subStepLen)
178 : {
179 0 : if (roundIdx == 0) {
180 : // 给本次卡的数据
181 : SendMemBlock sendData0;
182 0 : CHK_RET(CalcSendInfo(userRank_, myMinor_, times, subStepLen, sendData0));
183 0 : DeviceMem src0 = userInput_.range(sendData0.userInOffset, sendData0.sendLen);
184 0 : DeviceMem dst0 = cclInMem_.range(sendData0.cclDstOffset, sendData0.sendLen);
185 0 : HCCL_DEBUG(
186 : "[AlltoAllVFor310P] rank[%u] localCopy to cclIn, src.Offset [%llu], dst.Offset[%llu], len[%llu]", userRank_,
187 : sendData0.userInOffset, sendData0.cclDstOffset, sendData0.sendLen);
188 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst0, src0, mainStream_));
189 0 : }
190 :
191 : // 给右次卡的数据
192 : SendMemBlock sendData1;
193 0 : CHK_RET(CalcSendInfo(userRank_, rightMinor_, times, subStepLen, sendData1));
194 0 : DeviceMem src1 = userInput_.range(sendData1.userInOffset, sendData1.sendLen);
195 0 : DeviceMem dst1 = cclInMem_.range(cclBlockSize_ + sendData1.cclDstOffset, sendData1.sendLen);
196 0 : HCCL_DEBUG(
197 : "[AlltoAllVFor310P] rank[%u] localCopy to cclIn, src.Offset [%llu], dst.Offset[%llu], len[%llu]", userRank_,
198 : sendData1.userInOffset, cclBlockSize_ + sendData1.cclDstOffset, sendData1.sendLen);
199 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst1, src1, mainStream_));
200 :
201 0 : HCCL_DEBUG("[AlltoAllVFor310P] MainFirstLocalCopy finish.");
202 0 : return HCCL_SUCCESS;
203 0 : }
204 :
205 0 : HcclResult AlltoAllVFor310P::MinorFirstLocalCopy(const u32 times, const u32 roundIdx, const u64 subStepLen)
206 : {
207 : (void)roundIdx;
208 : // 给右次卡的数据
209 : SendMemBlock sendData;
210 0 : CHK_RET(CalcSendInfo(userRank_, rightMinor_, times, subStepLen, sendData));
211 0 : DeviceMem src = userInput_.range(sendData.userInOffset, sendData.sendLen);
212 0 : DeviceMem dst = cclInMem_.range(sendData.cclDstOffset, sendData.sendLen);
213 0 : HCCL_DEBUG(
214 : "[AlltoAllVFor310P] rank[%u] localCopy to cclIn, src.Offset [%llu], dst.Offset[%llu], len[%llu]", userRank_,
215 : sendData.userInOffset, sendData.cclDstOffset, sendData.sendLen);
216 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, mainStream_));
217 :
218 0 : return HCCL_SUCCESS;
219 0 : }
220 :
221 0 : HcclResult AlltoAllVFor310P::RunAlltoAllVFor310P()
222 : {
223 0 : u32 roundNum = userRankSize_ == DUO_RANK_NUM ? 1 : DUO_RANK_NUM - 1;
224 0 : u64 remainLen = CalcMaxSendLength();
225 0 : u64 subStepLen = std::min(remainLen, maxSizePerLoop_);
226 0 : u32 totalTimes = (remainLen + subStepLen - 1) / subStepLen;
227 0 : HCCL_INFO(
228 : "[AlltoAllVFor310P] roundNum[%u], maxLength[%llu], subStepLen[%llu], totalTimes[%u]", roundNum, remainLen,
229 : subStepLen, totalTimes);
230 0 : for (u32 times = 0; times < totalTimes && remainLen > 0; times++) {
231 0 : subStepLen = std::min(remainLen, maxSizePerLoop_);
232 0 : for (u32 roundIdx = 0; roundIdx < roundNum; roundIdx++) {
233 0 : SetNeighborRanks(roundIdx);
234 0 : for (u32 stepIdx = 0; stepIdx < STEP_NUM; stepIdx++) {
235 0 : CHK_RET(UpdateSendRecvRankInfo(roundIdx, stepIdx));
236 0 : CHK_RET(RunSendRecvBuffer(times, roundIdx, stepIdx, subStepLen));
237 : }
238 0 : HCCL_INFO("[AlltoAllVFor310P] Round[%u] finish.", roundIdx);
239 : }
240 0 : remainLen = remainLen - subStepLen;
241 0 : HCCL_INFO("[AlltoAllVFor310P] Times[%u] finish.", times);
242 : }
243 0 : return HCCL_SUCCESS;
244 : }
245 :
246 0 : void AlltoAllVFor310P::SetNeighborRanks(const u32 roundIdx)
247 : {
248 0 : if (userRank_ % COMPUTE_CONST == 0) {
249 0 : rightMain_ = ((userRank_ + COMPUTE_CONST * (roundIdx + 1))) % userRankSize_;
250 0 : rightMinor_ = ((userRank_ + MAX_RANK_GAP * (roundIdx + 1))) % userRankSize_;
251 0 : leftMain_ = ((userRank_ - COMPUTE_CONST * (roundIdx + 1)) + userRankSize_) % userRankSize_;
252 0 : leftMinor_ = ((userRank_ - 1 * (roundIdx + 1)) + userRankSize_) % userRankSize_;
253 : } else {
254 0 : rightMain_ = ((userRank_ + 1 * (roundIdx + 1))) % userRankSize_;
255 0 : rightMinor_ = ((userRank_ + COMPUTE_CONST * (roundIdx + 1))) % userRankSize_;
256 0 : leftMain_ = ((userRank_ - MAX_RANK_GAP * (roundIdx + 1)) + userRankSize_) % userRankSize_;
257 0 : leftMinor_ = ((userRank_ - COMPUTE_CONST * (roundIdx + 1)) + userRankSize_) % userRankSize_;
258 : }
259 0 : HCCL_DEBUG("[AlltoAllVFor310P] SetNeighborRanks finish.");
260 0 : }
261 :
262 0 : HcclResult AlltoAllVFor310P::UpdateSendRecvRankInfo(const u32 roundIdx, const u32 stepIdx)
263 : {
264 0 : sendRecvRankInfo_.clear();
265 0 : if (stepIdx <= THIRD_STEP && mainRank_) {
266 0 : sendRecvRankInfo_.push_back(std::make_pair(rightMain_, leftMain_)); // first send, second recv
267 0 : sendRecvRankInfo_.push_back(std::make_pair(myMinor_, myMinor_));
268 0 : } else if (stepIdx <= THIRD_STEP && minorRank_) {
269 0 : sendRecvRankInfo_.push_back(std::make_pair(myMain_, myMain_));
270 0 : } else if (stepIdx > THIRD_STEP && mainRank_) {
271 0 : sendRecvRankInfo_.push_back(std::make_pair(rightMain_, leftMain_));
272 : }
273 0 : HCCL_DEBUG(
274 : "[AlltoAllVFor310P][UpdateSendRecvRankInfo] update send/recv rank finished, roundIdx[%u], stepIdx[%u]",
275 : roundIdx, stepIdx);
276 0 : return HCCL_SUCCESS;
277 : }
278 :
279 : HcclResult
280 0 : AlltoAllVFor310P::RunMainCommonSteps(const u32 times, const u32 roundIdx, const u32 stepIdx, const u64 subStepLen)
281 : {
282 0 : UpdateMainStepMemInfo(roundIdx, stepIdx);
283 0 : if (stepIdx == THIRD_STEP) {
284 : // 给其他主卡的数据拷到cclOut
285 : SendMemBlock sendData;
286 0 : CHK_RET(CalcSendInfo(userRank_, rightMain_, times, subStepLen, sendData));
287 0 : DeviceMem src = userInput_.range(sendData.userInOffset, sendData.sendLen);
288 0 : DeviceMem dst = cclOutMem_.range(sendData.cclDstOffset, sendData.sendLen);
289 0 : HCCL_DEBUG(
290 : "[AlltoAllVFor310P] rank[%u] localCopy to cclOut, src.Offset [%llu], dst.Offset[%llu], len[%llu]",
291 : userRank_, sendData.userInOffset, sendData.cclDstOffset, sendData.sendLen);
292 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, mainStream_));
293 0 : }
294 :
295 : RecvMemBlock recvData0;
296 0 : CHK_RET(CalcRecvInfo(mainStepInfo_.readMinor.first, mainStepInfo_.readMinor.second, times, subStepLen, recvData0));
297 0 : u64 dstOffset = 0;
298 0 : if (stepIdx == THIRD_STEP) {
299 0 : dstOffset = recvData0.userOutOffset;
300 : } else {
301 0 : dstOffset = cclBlockSize_ + recvData0.cclSrcOffset;
302 : }
303 0 : const LINK& readMinorTransport = links_[sendRecvRankInfo_[1].second];
304 0 : CHK_PTR_NULL(readMinorTransport);
305 0 : void* remMemPtr0 = nullptr;
306 0 : CHK_RET(readMinorTransport->GetRemoteMem(mainStepInfo_.srcMemType, &remMemPtr0));
307 0 : DeviceMem remoteMem0 = DeviceMem::create(static_cast<u8*>(remMemPtr0), memList_[mainStepInfo_.srcBuffId].size());
308 0 : DeviceMem src0 = remoteMem0.range(recvData0.cclSrcOffset, recvData0.recvLen);
309 0 : DeviceMem dst0 = memList_[mainStepInfo_.dstBuffId].range(dstOffset, recvData0.recvLen);
310 0 : HCCL_DEBUG(
311 : "[AlltoAllVFor310P] rank[%u] Memcpy to rank[%u], src.Offset [%llu], dst.Offset[%llu], len[%llu]",
312 : sendRecvRankInfo_[1].second, userRank_, recvData0.cclSrcOffset, dstOffset, recvData0.recvLen);
313 0 : CHK_RET(HcclD2DMemcpyAsync(
314 : dispatcher_, dst0, src0, subStream_[1], readMinorTransport->GetRemoteRank(),
315 : readMinorTransport->GetLinkType()));
316 :
317 : RecvMemBlock recvData1;
318 0 : CHK_RET(CalcRecvInfo(mainStepInfo_.readMain.first, mainStepInfo_.readMain.second, times, subStepLen, recvData1));
319 0 : if (stepIdx == THIRD_STEP) {
320 0 : dstOffset = recvData1.userOutOffset;
321 : } else {
322 0 : dstOffset = recvData1.cclSrcOffset;
323 : }
324 0 : const LINK& readMainTransport = links_[sendRecvRankInfo_[0].second];
325 0 : CHK_PTR_NULL(readMainTransport);
326 0 : void* remMemPtr1 = nullptr;
327 0 : CHK_RET(readMainTransport->GetRemoteMem(mainStepInfo_.srcMemType, &remMemPtr1));
328 0 : DeviceMem remoteMem1 = DeviceMem::create(static_cast<u8*>(remMemPtr1), memList_[mainStepInfo_.srcBuffId].size());
329 0 : DeviceMem src1 = remoteMem1.range(cclBlockSize_ + recvData1.cclSrcOffset, recvData1.recvLen);
330 0 : DeviceMem dst1 = memList_[mainStepInfo_.dstBuffId].range(dstOffset, recvData1.recvLen);
331 0 : HCCL_DEBUG(
332 : "[AlltoAllVFor310P] rank[%u] Memcpy to rank[%u], src.Offset [%llu], dst.Offset[%llu], len[%llu]",
333 : sendRecvRankInfo_[0].second, userRank_, cclBlockSize_ + recvData1.cclSrcOffset, dstOffset, recvData1.recvLen);
334 0 : CHK_RET(HcclD2DMemcpyAsync(
335 : dispatcher_, dst1, src1, subStream_[0], readMainTransport->GetRemoteRank(), readMainTransport->GetLinkType()));
336 0 : HCCL_DEBUG("[AlltoAllVFor310P] RunMainStep %u finish.", stepIdx);
337 0 : return HCCL_SUCCESS;
338 0 : }
339 :
340 0 : void AlltoAllVFor310P::UpdateMainStepMemInfo(const u32 roundIdx, const u32 stepIdx)
341 : {
342 : (void)roundIdx;
343 0 : if (stepIdx % COMPUTE_CONST == 1) {
344 0 : mainStepInfo_.srcBuffId = 0;
345 0 : mainStepInfo_.srcMemType = UserMemType::INPUT_MEM;
346 : } else {
347 0 : mainStepInfo_.srcBuffId = 1;
348 0 : mainStepInfo_.srcMemType = UserMemType::OUTPUT_MEM;
349 : }
350 0 : if (stepIdx == 1) {
351 0 : mainStepInfo_.dstBuffId = 1;
352 0 : mainStepInfo_.readMain = std::make_pair(leftMain_, myMinor_);
353 0 : mainStepInfo_.readMinor = std::make_pair(myMinor_, rightMinor_);
354 0 : } else if (stepIdx == COMPUTE_CONST) {
355 0 : mainStepInfo_.dstBuffId = 0;
356 0 : mainStepInfo_.readMain = std::make_pair(leftMinor_, myMinor_);
357 0 : mainStepInfo_.readMinor = std::make_pair(myMinor_, rightMain_);
358 : } else {
359 0 : mainStepInfo_.dstBuffId = COMPUTE_CONST;
360 0 : mainStepInfo_.readMain = std::make_pair(leftMinor_, userRank_);
361 0 : mainStepInfo_.readMinor = std::make_pair(myMinor_, userRank_);
362 : }
363 0 : }
364 :
365 0 : HcclResult AlltoAllVFor310P::RunMainStep4(const u32 times, const u64 subStepLen)
366 : {
367 : // 拷本卡的数据
368 : SendMemBlock sendDataLocal;
369 0 : CHK_RET(CalcSendInfo(userRank_, userRank_, times, subStepLen, sendDataLocal));
370 : RecvMemBlock recvDataLocal;
371 0 : CHK_RET(CalcRecvInfo(userRank_, userRank_, times, subStepLen, recvDataLocal));
372 :
373 0 : DeviceMem src0 = userInput_.range(sendDataLocal.userInOffset, sendDataLocal.sendLen);
374 0 : DeviceMem dst0 = userOutput_.range(recvDataLocal.userOutOffset, recvDataLocal.recvLen);
375 0 : HCCL_DEBUG(
376 : "[AlltoAllVFor310P] rank[%u] localCopy to userOut, src.Offset [%llu], dst.Offset[%llu], len[%llu]", userRank_,
377 : sendDataLocal.userInOffset, recvDataLocal.userOutOffset, recvDataLocal.recvLen);
378 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst0, src0, mainStream_));
379 :
380 : // 读其他主卡的数据到usrOut
381 : RecvMemBlock recvData;
382 0 : CHK_RET(CalcRecvInfo(leftMain_, userRank_, times, subStepLen, recvData));
383 0 : const LINK& readMainTransport = links_[sendRecvRankInfo_[0].second];
384 0 : CHK_PTR_NULL(readMainTransport);
385 0 : void* remMemPtr = nullptr;
386 0 : CHK_RET(readMainTransport->GetRemoteMem(UserMemType::OUTPUT_MEM, &remMemPtr));
387 0 : DeviceMem remoteCCLInMem = DeviceMem::create(static_cast<u8*>(remMemPtr), cclOutMem_.size());
388 0 : DeviceMem src1 = remoteCCLInMem.range(recvData.cclSrcOffset, recvData.recvLen);
389 0 : DeviceMem dst1 = userOutput_.range(recvData.userOutOffset, recvData.recvLen);
390 0 : HCCL_DEBUG(
391 : "[AlltoAllVFor310P] rank[%u] Memcpy to rank[%u] userOut, src.Offset [%llu], dst.Offset[%llu], len[%llu]",
392 : sendRecvRankInfo_[0].second, userRank_, recvData.cclSrcOffset, recvData.userOutOffset, recvData.recvLen);
393 0 : CHK_RET(HcclD2DMemcpyAsync(
394 : dispatcher_, dst1, src1, subStream_[0], readMainTransport->GetRemoteRank(), readMainTransport->GetLinkType()));
395 0 : HCCL_DEBUG("[AlltoAllVFor310P] RunMainStep 4 finish.");
396 0 : return HCCL_SUCCESS;
397 0 : }
398 :
399 0 : u64 AlltoAllVFor310P::CalcMaxSendLength()
400 : {
401 0 : u64 maxSendDataLen = 0;
402 0 : for (u32 i = 0; i < allMeshAggregationSendRecvInfoPtr_->size(); i++) {
403 0 : for (u32 j = 0; j < (*allMeshAggregationSendRecvInfoPtr_)[i].sendLength.size(); j++) {
404 0 : u64 sendLength = (*allMeshAggregationSendRecvInfoPtr_)[i].sendLength[j];
405 0 : maxSendDataLen = std::max(maxSendDataLen, sendLength);
406 : }
407 : }
408 0 : HCCL_DEBUG("[AlltoAllVFor310P][CalcMaxSendLength] maxSendDataLen[%llu]", maxSendDataLen);
409 0 : return maxSendDataLen;
410 : }
411 :
412 : HcclResult
413 0 : AlltoAllVFor310P::RunMainSendRecvBuffer(const u32 times, const u32 roundIdx, const u32 stepIdx, const u64 subStepLen)
414 : {
415 : // 读次卡的都用从流1,读主卡的都用从流0
416 0 : if (stepIdx == 0) {
417 0 : CHK_RET(MainFirstLocalCopy(times, roundIdx, subStepLen));
418 0 : } else if (stepIdx <= THIRD_STEP) {
419 0 : CHK_RET(NotifySubStreamStart());
420 0 : const LINK& rightMainTransport = links_[rightMain_];
421 0 : const LINK& leftMainTransport = links_[leftMain_];
422 0 : const LINK& minorTransport = links_[myMinor_];
423 0 : CHK_RET(rightMainTransport->TxAck(subStream_[0]));
424 0 : CHK_RET(leftMainTransport->RxAck(subStream_[0]));
425 0 : CHK_RET(minorTransport->TxAck(subStream_[1]));
426 0 : CHK_RET(minorTransport->RxAck(subStream_[1]));
427 0 : CHK_RET(RunMainCommonSteps(times, roundIdx, stepIdx, subStepLen));
428 0 : CHK_RET(leftMainTransport->TxDataSignal(subStream_[0]));
429 0 : CHK_RET(rightMainTransport->RxDataSignal(subStream_[0]));
430 0 : CHK_RET(minorTransport->TxDataSignal(subStream_[1]));
431 0 : CHK_RET(minorTransport->RxDataSignal(subStream_[1]));
432 0 : CHK_RET(WaitSubStreamFinish());
433 : } else {
434 0 : CHK_RET(NotifySubStreamStart());
435 0 : const LINK& rightMainTransport = links_[rightMain_];
436 0 : const LINK& leftMainTransport = links_[leftMain_];
437 0 : CHK_RET(rightMainTransport->TxAck(subStream_[0]));
438 0 : CHK_RET(leftMainTransport->RxAck(subStream_[0]));
439 0 : CHK_RET(RunMainStep4(times, subStepLen));
440 0 : CHK_RET(leftMainTransport->TxDataSignal(subStream_[0]));
441 0 : CHK_RET(rightMainTransport->RxDataSignal(subStream_[0]));
442 0 : CHK_RET(WaitSubStreamFinish());
443 : }
444 0 : return HCCL_SUCCESS;
445 : }
446 :
447 : HcclResult
448 0 : AlltoAllVFor310P::RunMinorSendRecvBuffer(const u32 times, const u32 roundIdx, const u32 stepIdx, const u64 subStepLen)
449 : {
450 0 : if (stepIdx == 0) {
451 0 : CHK_RET(MinorFirstLocalCopy(times, roundIdx, subStepLen));
452 : // 主流告诉从流已拷贝完
453 0 : } else if (stepIdx <= THIRD_STEP) {
454 0 : CHK_RET(NotifySubStreamStart());
455 0 : const LINK& mainTransport = links_[myMain_]; // 次die读主die 0
456 0 : CHK_RET(mainTransport->TxAck(subStream_[0]));
457 0 : CHK_RET(mainTransport->RxAck(subStream_[0]));
458 0 : CHK_RET(RunMinorCommonSteps(times, roundIdx, stepIdx, subStepLen));
459 0 : CHK_RET(mainTransport->TxDataSignal(subStream_[0]));
460 0 : CHK_RET(mainTransport->RxDataSignal(subStream_[0]));
461 0 : CHK_RET(WaitSubStreamFinish());
462 : } else {
463 0 : CHK_RET(NotifySubStreamStart());
464 0 : CHK_RET(RunMinorStep4(times, subStepLen));
465 0 : CHK_RET(WaitSubStreamFinish());
466 : }
467 0 : return HCCL_SUCCESS;
468 : }
469 :
470 0 : void AlltoAllVFor310P::UpdateMinorStepMemInfo(const u32 roundIdx, const u32 stepIdx)
471 : {
472 : (void)roundIdx;
473 0 : if (stepIdx % COMPUTE_CONST == 1) {
474 0 : minorStepInfo_.srcBuffId = 0; // 读主卡的src buffer
475 0 : minorStepInfo_.srcMemType = UserMemType::INPUT_MEM;
476 0 : minorStepInfo_.dstBuffId = 1; // 本地拷贝的dst buffer
477 : } else {
478 0 : minorStepInfo_.srcBuffId = 1;
479 0 : minorStepInfo_.srcMemType = UserMemType::OUTPUT_MEM;
480 0 : minorStepInfo_.dstBuffId = 0;
481 : }
482 0 : if (stepIdx == 1) {
483 0 : minorStepInfo_.readMain = std::make_pair(myMain_, userRank_);
484 0 : minorStepInfo_.readMinor = std::make_pair(userRank_, rightMain_); // 本地拷贝的数据的src/dst rank
485 0 : } else if (stepIdx == COMPUTE_CONST) {
486 0 : minorStepInfo_.readMain = std::make_pair(leftMain_, userRank_);
487 0 : minorStepInfo_.readMinor = std::make_pair(userRank_, myMain_);
488 : } else {
489 0 : minorStepInfo_.readMain = std::make_pair(leftMinor_, userRank_);
490 : }
491 0 : }
492 :
493 : HcclResult
494 0 : AlltoAllVFor310P::RunMinorCommonSteps(const u32 times, const u32 roundIdx, const u32 stepIdx, const u64 subStepLen)
495 : {
496 0 : UpdateMinorStepMemInfo(roundIdx, stepIdx);
497 0 : if (stepIdx != THIRD_STEP) {
498 : SendMemBlock sendData;
499 0 : CHK_RET(
500 : CalcSendInfo(minorStepInfo_.readMinor.first, minorStepInfo_.readMinor.second, times, subStepLen, sendData));
501 0 : DeviceMem src = userInput_.range(sendData.userInOffset, sendData.sendLen);
502 0 : DeviceMem dst = memList_[minorStepInfo_.dstBuffId].range(sendData.cclDstOffset, sendData.sendLen);
503 0 : HCCL_DEBUG(
504 : "[AlltoAllVFor310P] rank[%u] localcopy, src.Offset [%llu], dst.Offset[%llu], len[%llu]", userRank_,
505 : sendData.userInOffset, sendData.cclDstOffset, sendData.sendLen);
506 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, mainStream_));
507 0 : }
508 :
509 : RecvMemBlock recvData;
510 0 : CHK_RET(CalcRecvInfo(minorStepInfo_.readMain.first, minorStepInfo_.readMain.second, times, subStepLen, recvData));
511 0 : const LINK& readMainTransport = links_[sendRecvRankInfo_[0].second];
512 0 : CHK_PTR_NULL(readMainTransport);
513 0 : void* remMemPtr = nullptr;
514 0 : CHK_RET(readMainTransport->GetRemoteMem(minorStepInfo_.srcMemType, &remMemPtr));
515 0 : DeviceMem remoteMem = DeviceMem::create(static_cast<u8*>(remMemPtr), memList_[minorStepInfo_.srcBuffId].size());
516 0 : DeviceMem src0 = remoteMem.range(recvData.cclSrcOffset, recvData.recvLen);
517 0 : DeviceMem dst0 = userOutput_.range(recvData.userOutOffset, recvData.recvLen);
518 0 : HCCL_DEBUG(
519 : "[AlltoAllVFor310P] rank[%u] Memcpy to rank[%u] userOut, src.Offset [%llu], dst.Offset[%llu], len[%llu]",
520 : sendRecvRankInfo_[0].second, userRank_, recvData.cclSrcOffset, recvData.userOutOffset, recvData.recvLen);
521 0 : CHK_RET(HcclD2DMemcpyAsync(
522 : dispatcher_, dst0, src0, subStream_[0], readMainTransport->GetRemoteRank(), readMainTransport->GetLinkType()));
523 0 : HCCL_DEBUG("[AlltoAllVFor310P] RunMinorStep %u finish.", stepIdx);
524 :
525 0 : return HCCL_SUCCESS;
526 0 : }
527 :
528 0 : HcclResult AlltoAllVFor310P::RunMinorStep4(const u32 times, const u64 subStepLen)
529 : {
530 : SendMemBlock sendDataLocal;
531 0 : CHK_RET(CalcSendInfo(userRank_, userRank_, times, subStepLen, sendDataLocal));
532 : RecvMemBlock recvDataLocal;
533 0 : CHK_RET(CalcRecvInfo(userRank_, userRank_, times, subStepLen, recvDataLocal));
534 :
535 0 : DeviceMem src = userInput_.range(sendDataLocal.userInOffset, sendDataLocal.sendLen);
536 0 : DeviceMem dst = userOutput_.range(recvDataLocal.userOutOffset, recvDataLocal.recvLen);
537 0 : HCCL_DEBUG(
538 : "[AlltoAllVFor310P] rank[%u] localCopy to userOut, src.Offset [%llu], dst.Offset[%llu], len[%llu]", userRank_,
539 : sendDataLocal.userInOffset, recvDataLocal.userOutOffset, recvDataLocal.recvLen);
540 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, mainStream_));
541 0 : HCCL_DEBUG("[AlltoAllVFor310P] RunMinorStep 4 finish.");
542 0 : return HCCL_SUCCESS;
543 0 : }
544 :
545 : HcclResult
546 0 : AlltoAllVFor310P::RunSendRecvBuffer(const u32 times, const u32 roundIdx, const u32 stepIdx, const u64 subStepLen)
547 : {
548 : // 读次卡的都用从流0,读主卡的都用从流1
549 0 : HCCL_INFO(
550 : "[AlltoAllVFor310P] RunSendRecvBuffer start, times[%u], roundIdx[%u], stepIdx[%u]", times, roundIdx, stepIdx);
551 0 : if (mainRank_) {
552 0 : CHK_RET(RunMainSendRecvBuffer(times, roundIdx, stepIdx, subStepLen));
553 : } else {
554 0 : CHK_RET(RunMinorSendRecvBuffer(times, roundIdx, stepIdx, subStepLen));
555 : }
556 0 : return HCCL_SUCCESS;
557 : }
558 :
559 0 : HcclResult AlltoAllVFor310P::RunAsync()
560 : {
561 0 : HcclOpMetaInfoDef opMeta = HcclOpMetaInfo::GetOneForAllToAllV(CopyPattern::ZCOPY, cclInMem_.size(), true);
562 0 : CHK_RET(InitTask(dispatcher_, mainStream_, opMeta.isEnableCache, opMeta.GetCacheKey()));
563 0 : CHK_RET(RunAlltoAllVFor310P());
564 0 : CHK_RET(LaunchTaskExtend(dispatcher_, mainStream_, subStream_));
565 0 : HCCL_INFO("[AlltoAllVFor310P][RunAsync] finished");
566 0 : return HCCL_SUCCESS;
567 : }
568 : REGISTER_TEMPLATE(TemplateType::TEMPLATE_ALL_2_ALL_V_FOR310P, AlltoAllVFor310P);
569 : // namespace hccl
570 : } // namespace hccl
|