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_retry_executor.h"
12 : namespace hccl {
13 : constexpr u32 PAIRSIZE_TWO = 2;
14 :
15 0 : CollBatchSendRecvRetryExecutor::CollBatchSendRecvRetryExecutor(
16 0 : const HcclDispatcher dispatcher, std::unique_ptr<TopoMatcher>& topoMatcher)
17 0 : : CollBatchSendRecvExecutor(dispatcher, topoMatcher)
18 0 : {}
19 :
20 0 : HcclResult CollBatchSendRecvRetryExecutor::CreatePairWiseList(HcclSendRecvItem* sendRecvInfo, u32 itemNum)
21 : {
22 0 : HCCL_INFO("[CollBatchSendRecvRetryExecutor][GetPairWiseList] Start sort the batchSendRecv tasklist.");
23 0 : CHK_PTR_NULL(sendRecvInfo);
24 :
25 0 : for (u32 i = 0; i < itemNum; i++) {
26 0 : HCCL_INFO(
27 : "[CollBatchSendRecvRetryExecutor][GetPairWiseList] index is %u, itemNum is %u, localRankID is %u, "
28 : "remoteRank is %u, sendRecvType is %u, rankSize is %u.",
29 : i, itemNum, topoAttr_.userRank, sendRecvInfo->remoteRank, static_cast<u32>(sendRecvInfo->sendRecvType),
30 : topoAttr_.userRankSize);
31 0 : CHK_PTR_NULL(sendRecvInfo->buf);
32 :
33 0 : if (sendRecvInfo->sendRecvType == HcclSendRecvType::HCCL_SEND) {
34 0 : sendDeque_.push_back(sendRecvInfo);
35 0 : } else if (sendRecvInfo->sendRecvType == HcclSendRecvType::HCCL_RECV) {
36 0 : recvDeque_.push_back(sendRecvInfo);
37 : } else {
38 0 : HCCL_ERROR(
39 : "[CollBatchSendRecvRetryExecutor][GetPairWiseList] sendRecvType wrong sendrecvType is %d, "
40 : "rankID is %u, remoteRank is %u.",
41 : sendRecvInfo->sendRecvType, topoAttr_.userRank, sendRecvInfo->remoteRank);
42 0 : return HCCL_E_PARA;
43 : }
44 0 : sendRecvInfo++;
45 : }
46 : /* 此处的排序逻辑(pair-wise算法):
47 : 1.sendDeque元素顺序是:先放remoteRank号小于等于root rank的第一个任务,依次减小(循环索引)直至放完
48 : 2.recvDeque元素顺序是:先放remoteRank号大于等于root rank的第一个任务,依次增大(循环索引)直至放完
49 : */
50 0 : auto sendCompare = [this](HcclSendRecvItem* a, HcclSendRecvItem* b) {
51 0 : u32 aFlag = (a->remoteRank <= topoAttr_.userRank) ? (a->remoteRank + topoAttr_.userRankSize) : a->remoteRank;
52 0 : u32 bFlag = (b->remoteRank <= topoAttr_.userRank) ? (b->remoteRank + topoAttr_.userRankSize) : b->remoteRank;
53 0 : return aFlag > bFlag;
54 0 : };
55 :
56 0 : auto recvCompare = [this](HcclSendRecvItem* a, HcclSendRecvItem* b) {
57 0 : u32 aFlag = (a->remoteRank < topoAttr_.userRank) ? (a->remoteRank + topoAttr_.userRankSize) : a->remoteRank;
58 0 : u32 bFlag = (b->remoteRank < topoAttr_.userRank) ? (b->remoteRank + topoAttr_.userRankSize) : b->remoteRank;
59 0 : return aFlag < bFlag;
60 0 : };
61 :
62 0 : std::sort(sendDeque_.begin(), sendDeque_.end(), sendCompare);
63 0 : std::sort(recvDeque_.begin(), recvDeque_.end(), recvCompare);
64 :
65 : // 生成SendRecvPair
66 0 : u32 pairNum = std::max(sendDeque_.size(), recvDeque_.size());
67 0 : for (u32 pairIndex = 0; pairIndex < pairNum; pairIndex++) {
68 0 : std::vector<HcclSendRecvItem*> sendRecvPair;
69 0 : if (sendDeque_.size() > pairIndex) {
70 0 : sendRecvPair.push_back(sendDeque_[pairIndex]);
71 : }
72 0 : if (recvDeque_.size() > pairIndex) {
73 0 : sendRecvPair.push_back(recvDeque_[pairIndex]);
74 : }
75 0 : sendRecvPairList_.push_back(sendRecvPair);
76 0 : }
77 0 : HCCL_INFO("[CollBatchSendRecvRetryExecutor][GetPairWiseList] End sort the batchSendRecv tasklist.");
78 0 : return HCCL_SUCCESS;
79 : }
80 :
81 : HcclResult
82 0 : CollBatchSendRecvRetryExecutor::GetPairWiseList(std::vector<std::vector<HcclSendRecvItem*>>& sendRecvPairList)
83 : {
84 0 : sendRecvPairList = sendRecvPairList_;
85 0 : return HCCL_SUCCESS;
86 : }
87 :
88 0 : HcclResult CollBatchSendRecvRetryExecutor::GetPairWiseList(HcclSendRecvItem* sendRecvInfo, u32 itemNum)
89 : {
90 0 : return CollBatchSendRecvExecutor::GetPairWiseList(sendRecvInfo, itemNum);
91 : }
92 :
93 0 : HcclResult CollBatchSendRecvRetryExecutor::CheckSendRecvPair(const std::vector<HcclSendRecvItem*>& sendRecvPair)
94 : {
95 0 : if (sendRecvPair.empty()) {
96 0 : HCCL_ERROR("[CollBatchSendRecvRetryExecutor] please check the pair list.");
97 0 : return HCCL_E_PARA;
98 : }
99 0 : if (sendRecvPair.size() == 1 && sendRecvPair[0]->remoteRank == topoAttr_.userRank) {
100 0 : HCCL_ERROR("[CollBatchSendRecvRetryExecutor] SendTask and Recv Task to rank itself do not match,"
101 : "please check the task list.");
102 0 : return HCCL_E_PARA;
103 : }
104 0 : return HCCL_SUCCESS;
105 : }
106 :
107 0 : HcclResult CollBatchSendRecvRetryExecutor::Orchestrate(OpParam& param, AlgResourceResponse& algResource)
108 : {
109 0 : HcclUs startut = TIME_NOW();
110 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[CollBatchSendRecvRetryExecutor] batchsendrecv retry starts.");
111 0 : algResResp_ = &algResource;
112 0 : CHK_RET(CheckCommSize(COMM_COMBINE_ORDER, COMM_SIZE_TWO));
113 :
114 : // 校验当前sendRecvPair
115 0 : std::vector<HcclSendRecvItem*> sendRecvPair;
116 0 : if (param.BatchSendRecvDataDes.curIterNum < sendRecvPairList_.size()) {
117 0 : sendRecvPair = sendRecvPairList_[param.BatchSendRecvDataDes.curIterNum];
118 : } else {
119 0 : HCCL_ERROR(
120 : "[CollBatchSendRecvRetryExecutor] the curIterNum[%u] is out of range[0, %zu].",
121 : param.BatchSendRecvDataDes.curIterNum, sendRecvPairList_.size());
122 0 : return HCCL_E_PARA;
123 : }
124 0 : CHK_RET(CheckSendRecvPair(sendRecvPair));
125 :
126 : // 自发自收场景
127 0 : if (sendRecvPair.size() == PAIRSIZE_TWO && sendRecvPair[0]->remoteRank == topoAttr_.userRank
128 0 : && sendRecvPair[1]->remoteRank == topoAttr_.userRank) {
129 0 : if (sendRecvPair[0]->count == sendRecvPair[1]->count
130 0 : && sendRecvPair[0]->dataType == sendRecvPair[1]->dataType) {
131 0 : u64 dataSize = sendRecvPair[0]->count * SIZE_TABLE[sendRecvPair[0]->dataType];
132 0 : DeviceMem inUserMem = DeviceMem::create(static_cast<u8*>(sendRecvPair[0]->buf), dataSize);
133 0 : DeviceMem outUserMem = DeviceMem::create(static_cast<u8*>(sendRecvPair[1]->buf), dataSize);
134 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, outUserMem, inUserMem, param.stream));
135 0 : return HCCL_SUCCESS;
136 0 : } else {
137 0 : HCCL_ERROR("[HcclBatchSendRecvRetry] Send task and recv task to self : data size do not equal, please"
138 : "check the task list.");
139 0 : return HCCL_E_PARA;
140 : }
141 : }
142 :
143 : // 重执行正常执行场景,前后需和控制流做同步
144 0 : if (param.BatchSendRecvDataDes.curMode == BatchSendRecvCurMode::SEND_RECV) {
145 0 : HCCL_INFO("[BatchSendRecv] Stream sync: main stream record, subStream wait.");
146 0 : CHK_RET(LocalNotify::Post(param.stream, dispatcher_, algResResp_->notifiesAux[STREAM_INDEX_0], PROF_STAGE_0));
147 0 : CHK_RET(LocalNotify::Wait(
148 : algResResp_->slaveStreams[STREAM_INDEX_0], dispatcher_, algResResp_->notifiesAux[STREAM_INDEX_0],
149 : PROF_STAGE_0));
150 0 : CHK_RET(LocalNotify::Post(param.stream, dispatcher_, algResResp_->notifiesAux[STREAM_INDEX_1], PROF_STAGE_1));
151 0 : CHK_RET(LocalNotify::Wait(
152 : algResResp_->slaveStreams[STREAM_INDEX_1], dispatcher_, algResResp_->notifiesAux[STREAM_INDEX_1],
153 : PROF_STAGE_1));
154 : }
155 : // run sendrecv
156 0 : CHK_RET(RunLoop(param, algResource, sendRecvPair));
157 :
158 0 : if (param.BatchSendRecvDataDes.curMode == BatchSendRecvCurMode::SEND_RECV) {
159 0 : HCCL_INFO("[BatchSendRecv] Stream sync: subStream record, main stream wait.");
160 0 : CHK_RET(LocalNotify::Post(
161 : algResResp_->slaveStreams[STREAM_INDEX_0], dispatcher_, algResResp_->notifiesMain[STREAM_INDEX_0],
162 : PROF_STAGE_0));
163 0 : CHK_RET(LocalNotify::Wait(param.stream, dispatcher_, algResResp_->notifiesMain[STREAM_INDEX_0], PROF_STAGE_0));
164 0 : CHK_RET(LocalNotify::Post(
165 : algResResp_->slaveStreams[STREAM_INDEX_1], dispatcher_, algResResp_->notifiesMain[STREAM_INDEX_1],
166 : PROF_STAGE_1));
167 0 : CHK_RET(LocalNotify::Wait(param.stream, dispatcher_, algResResp_->notifiesMain[STREAM_INDEX_1], PROF_STAGE_1));
168 0 : } else if (param.BatchSendRecvDataDes.curMode == BatchSendRecvCurMode::SEND) {
169 0 : CHK_RET(LocalNotify::Post(
170 : algResResp_->slaveStreams[STREAM_INDEX_0], dispatcher_, algResResp_->notifiesMain[STREAM_INDEX_0],
171 : PROF_STAGE_0));
172 0 : } else if (param.BatchSendRecvDataDes.curMode == BatchSendRecvCurMode::RECV) {
173 0 : CHK_RET(LocalNotify::Post(
174 : algResResp_->slaveStreams[STREAM_INDEX_1], dispatcher_, algResResp_->notifiesMain[STREAM_INDEX_1],
175 : PROF_STAGE_1));
176 : }
177 :
178 0 : CHK_RET(LaunchTaskExtend(dispatcher_, param.stream, algResResp_->slaveStreams));
179 0 : HCCL_INFO("[info][print] LaunchTaskExtend success.");
180 0 : HCCL_INFO(
181 : "tag[%s] BatchSendRecv Executor orchestrate success, take time [%lld]us.", param.tag.c_str(),
182 : DURATION_US(TIME_NOW() - startut));
183 0 : return HCCL_SUCCESS;
184 0 : }
185 :
186 0 : HcclResult CollBatchSendRecvRetryExecutor::RunLoop(
187 : OpParam& param, AlgResourceResponse& algRes, const std::vector<HcclSendRecvItem*>& sendRecvPair)
188 : {
189 : // 判断当前需执行的算子
190 0 : std::vector<HcclSendRecvItem*> curSendRecvPair;
191 0 : if (param.BatchSendRecvDataDes.curMode == BatchSendRecvCurMode::SEND) {
192 0 : curSendRecvPair.push_back(sendRecvPair[0]);
193 0 : } else if (param.BatchSendRecvDataDes.curMode == BatchSendRecvCurMode::RECV) {
194 0 : curSendRecvPair.push_back(sendRecvPair[sendRecvPair.size() - 1]);
195 : } else {
196 0 : curSendRecvPair = sendRecvPair;
197 : }
198 :
199 : // 执行当前需执行的算子
200 0 : for (const auto& itemPtr : curSendRecvPair) {
201 0 : if (static_cast<bool>(param.BatchSendRecvDataDes.isDirectRemoteRank[itemPtr->remoteRank])) {
202 : // device direct链路的任务会在host侧下发,此处需要跳过
203 0 : continue;
204 : }
205 0 : HCCL_INFO("[CollBatchSendRecvRetryExecutor][RunLoop] remoteRank %u", itemPtr->remoteRank);
206 0 : if (itemPtr->sendRecvType == HcclSendRecvType::HCCL_SEND) {
207 0 : CHK_RET(CalcSendSlices(algRes, itemPtr));
208 0 : } else if (itemPtr->sendRecvType == HcclSendRecvType::HCCL_RECV) {
209 0 : CHK_RET(CalcRecvSlices(algRes, itemPtr));
210 : } else {
211 0 : HCCL_ERROR("[CollBatchSendRecvRetryExecutor][RunLoop] sendRecvType is Wrong.");
212 0 : return HCCL_E_PARA;
213 : }
214 : }
215 :
216 0 : u32 loopInOnceLaunch = 0;
217 : // 每隔200个loop launch一次
218 0 : while (!sendDataSilces_.empty() || !recvDataSilces_.empty()) {
219 0 : if (!sendDataSilces_.empty()) {
220 0 : CHK_RET(ProcessSendDataSlice(algResResp_->slaveStreams[STREAM_INDEX_0], false, true));
221 0 : sendDataSilces_.pop_front();
222 : }
223 0 : if (!recvDataSilces_.empty()) {
224 0 : CHK_RET(ProcessRecvDataSlice(algResResp_->slaveStreams[STREAM_INDEX_1], true));
225 0 : recvDataSilces_.pop_front();
226 : }
227 0 : loopInOnceLaunch++;
228 0 : if (loopInOnceLaunch == MAX_LOOP_IN_ONCE_LAUNCH || (sendDataSilces_.empty() && recvDataSilces_.empty())) {
229 0 : CHK_RET(LaunchTaskExtend(dispatcher_, param.stream, algResResp_->slaveStreams));
230 0 : HCCL_INFO(
231 : "[BatchSendRecv] LaunchTaskExtend, unprocessed send slices[%u], recv slices[%u].",
232 : sendDataSilces_.size(), recvDataSilces_.size());
233 0 : loopInOnceLaunch = 0;
234 : }
235 : }
236 0 : return HCCL_SUCCESS;
237 0 : }
238 :
239 0 : HcclResult CollBatchSendRecvRetryExecutor::CalcStreamNum(u32& streamNum)
240 : {
241 0 : streamNum = LEVEL0_PLANE_NUM_IN_NPRING_DOUBLE;
242 0 : HCCL_INFO("[CollBatchSendRecvRetryExecutor][CalcScratchMemSize] tag_[%s], streamNum[%u].", tag_.c_str(), streamNum);
243 0 : return HCCL_SUCCESS;
244 : }
245 :
246 0 : HcclResult CollBatchSendRecvRetryExecutor::CalcSendSlices(AlgResourceResponse& algRes, HcclSendRecvItem* sendRecvItem)
247 : {
248 0 : HCCL_INFO(
249 : "[CollBatchSendRecvExecutor][CalcSendSlices] tag[%s], remoteRank[%u], buf[%p], count[%llu],"
250 : "dataType[%s], sendRecvType[%d].",
251 : tag_.c_str(), sendRecvItem->remoteRank, sendRecvItem->buf, sendRecvItem->count,
252 : GetDataTypeEnumStr(sendRecvItem->dataType).c_str(), sendRecvItem->sendRecvType);
253 0 : u8* curInputPtr = static_cast<u8*>(sendRecvItem->buf);
254 0 : CHK_PTR_NULL(curInputPtr);
255 0 : u32 unitSize = SIZE_TABLE[sendRecvItem->dataType];
256 0 : u64 maxCountPerLoop = CalcSendLoopMaxCount(const_cast<DeviceMem&>(algRes.cclInputMem), unitSize);
257 :
258 0 : for (u64 countLeft = sendRecvItem->count, curCount = 0, curOffset = 0; countLeft > 0; countLeft -= curCount) {
259 0 : curInputPtr += curOffset;
260 0 : curCount = (countLeft > maxCountPerLoop) ? maxCountPerLoop : countLeft;
261 0 : u64 curSize = curCount * unitSize; // 单位:字节
262 0 : sendDataSilces_.emplace_back(curInputPtr, curSize, sendRecvItem->remoteRank);
263 0 : HCCL_DEBUG(
264 : "[CollBatchSendRecvExecutor][CalcSendSlices] tag[%s], slice userAddr[%p], slice size[%llu].", tag_.c_str(),
265 : curInputPtr, curSize);
266 0 : curOffset = curSize;
267 : }
268 0 : return HCCL_SUCCESS;
269 : }
270 :
271 0 : HcclResult CollBatchSendRecvRetryExecutor::CalcRecvSlices(AlgResourceResponse& algRes, HcclSendRecvItem* sendRecvItem)
272 : {
273 0 : HCCL_INFO(
274 : "[CollBatchSendRecvRetryExecutor][CalcSendSlices] tag[%s], remoteRank[%u], buf[%p], count[%llu],"
275 : "dataType[%s], sendRecvType[%d].",
276 : tag_.c_str(), sendRecvItem->remoteRank, sendRecvItem->buf, sendRecvItem->count,
277 : GetDataTypeEnumStr(sendRecvItem->dataType).c_str(), sendRecvItem->sendRecvType);
278 0 : u8* curOutputPtr = static_cast<u8*>(sendRecvItem->buf);
279 0 : CHK_PTR_NULL(curOutputPtr);
280 0 : u32 unitSize = SIZE_TABLE[sendRecvItem->dataType];
281 0 : u64 maxCountPerLoop = CalcRecvLoopMaxCount(const_cast<DeviceMem&>(algRes.cclOutputMem), unitSize);
282 0 : HCCL_DEBUG("[CollBatchSendRecvRetryExecutor][CalcSendSlices]maxCountPerLoop is %llu", maxCountPerLoop);
283 :
284 0 : for (u64 countLeft = sendRecvItem->count, curCount = 0, curOffset = 0; countLeft > 0; countLeft -= curCount) {
285 0 : curOutputPtr += curOffset;
286 0 : curCount = (countLeft > maxCountPerLoop) ? maxCountPerLoop : countLeft;
287 0 : u64 curSize = curCount * unitSize; // 单位:字节
288 0 : recvDataSilces_.emplace_back(curOutputPtr, curSize, sendRecvItem->remoteRank);
289 0 : HCCL_DEBUG(
290 : "[CollBatchSendRecvRetryExecutor][CalcRecvSlices] tag[%s], slice userAddr[%p], slice size[%llu].",
291 : tag_.c_str(), curOutputPtr, curSize);
292 0 : curOffset = curSize;
293 : }
294 0 : return HCCL_SUCCESS;
295 : }
296 :
297 : REGISTER_EXEC("BatchSendRecvRetry", BatchSendRecvRetryExecutor, CollBatchSendRecvRetryExecutor);
298 : } // namespace hccl
|