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