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::CheckSendRecvPair(const std::vector<HcclSendRecvItem*>& sendRecvPair)
89 : {
90 0 : if (sendRecvPair.empty()) {
91 0 : HCCL_ERROR("[CollBatchSendRecvRetryExecutor] please check the pair list.");
92 0 : return HCCL_E_PARA;
93 : }
94 0 : if (sendRecvPair.size() == 1 && sendRecvPair[0]->remoteRank == topoAttr_.userRank) {
95 0 : HCCL_ERROR("[CollBatchSendRecvRetryExecutor] SendTask and Recv Task to rank itself do not match,"
96 : "please check the task list.");
97 0 : return HCCL_E_PARA;
98 : }
99 0 : return HCCL_SUCCESS;
100 : }
101 :
102 0 : HcclResult CollBatchSendRecvRetryExecutor::Orchestrate(OpParam& param, AlgResourceResponse& algResource)
103 : {
104 0 : HcclUs startut = TIME_NOW();
105 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[CollBatchSendRecvRetryExecutor] batchsendrecv retry starts.");
106 0 : algResResp_ = &algResource;
107 0 : CHK_RET(CheckCommSize(COMM_COMBINE_ORDER, COMM_SIZE_TWO));
108 :
109 : // 校验当前sendRecvPair
110 0 : std::vector<HcclSendRecvItem*> sendRecvPair;
111 0 : if (param.BatchSendRecvDataDes.curIterNum < sendRecvPairList_.size()) {
112 0 : sendRecvPair = sendRecvPairList_[param.BatchSendRecvDataDes.curIterNum];
113 : } else {
114 0 : HCCL_ERROR(
115 : "[CollBatchSendRecvRetryExecutor] the curIterNum[%u] is out of range[0, %zu].",
116 : param.BatchSendRecvDataDes.curIterNum, sendRecvPairList_.size());
117 0 : return HCCL_E_PARA;
118 : }
119 0 : CHK_RET(CheckSendRecvPair(sendRecvPair));
120 :
121 : // 自发自收场景
122 0 : if (sendRecvPair.size() == PAIRSIZE_TWO && sendRecvPair[0]->remoteRank == topoAttr_.userRank
123 0 : && sendRecvPair[1]->remoteRank == topoAttr_.userRank) {
124 0 : if (sendRecvPair[0]->count == sendRecvPair[1]->count
125 0 : && sendRecvPair[0]->dataType == sendRecvPair[1]->dataType) {
126 0 : u64 dataSize = sendRecvPair[0]->count * SIZE_TABLE[sendRecvPair[0]->dataType];
127 0 : DeviceMem inUserMem = DeviceMem::create(static_cast<u8*>(sendRecvPair[0]->buf), dataSize);
128 0 : DeviceMem outUserMem = DeviceMem::create(static_cast<u8*>(sendRecvPair[1]->buf), dataSize);
129 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, outUserMem, inUserMem, param.stream));
130 0 : return HCCL_SUCCESS;
131 0 : } else {
132 0 : HCCL_ERROR("[HcclBatchSendRecvRetry] Send task and recv task to self : data size do not equal, please"
133 : "check the task list.");
134 0 : return HCCL_E_PARA;
135 : }
136 : }
137 :
138 : // 重执行正常执行场景,前后需和控制流做同步
139 0 : if (param.BatchSendRecvDataDes.curMode == BatchSendRecvCurMode::SEND_RECV) {
140 0 : HCCL_INFO("[BatchSendRecv] Stream sync: main stream record, subStream wait.");
141 0 : CHK_RET(LocalNotify::Post(param.stream, dispatcher_, algResResp_->notifiesAux[STREAM_INDEX_0], PROF_STAGE_0));
142 0 : CHK_RET(LocalNotify::Wait(
143 : algResResp_->slaveStreams[STREAM_INDEX_0], dispatcher_, algResResp_->notifiesAux[STREAM_INDEX_0],
144 : PROF_STAGE_0));
145 0 : CHK_RET(LocalNotify::Post(param.stream, dispatcher_, algResResp_->notifiesAux[STREAM_INDEX_1], PROF_STAGE_1));
146 0 : CHK_RET(LocalNotify::Wait(
147 : algResResp_->slaveStreams[STREAM_INDEX_1], dispatcher_, algResResp_->notifiesAux[STREAM_INDEX_1],
148 : PROF_STAGE_1));
149 : }
150 : // run sendrecv
151 0 : CHK_RET(RunLoop(param, algResource, sendRecvPair));
152 :
153 0 : if (param.BatchSendRecvDataDes.curMode == BatchSendRecvCurMode::SEND_RECV) {
154 0 : HCCL_INFO("[BatchSendRecv] Stream sync: subStream record, main stream wait.");
155 0 : CHK_RET(LocalNotify::Post(
156 : algResResp_->slaveStreams[STREAM_INDEX_0], dispatcher_, algResResp_->notifiesMain[STREAM_INDEX_0],
157 : PROF_STAGE_0));
158 0 : CHK_RET(LocalNotify::Wait(param.stream, dispatcher_, algResResp_->notifiesMain[STREAM_INDEX_0], PROF_STAGE_0));
159 0 : CHK_RET(LocalNotify::Post(
160 : algResResp_->slaveStreams[STREAM_INDEX_1], dispatcher_, algResResp_->notifiesMain[STREAM_INDEX_1],
161 : PROF_STAGE_1));
162 0 : CHK_RET(LocalNotify::Wait(param.stream, dispatcher_, algResResp_->notifiesMain[STREAM_INDEX_1], PROF_STAGE_1));
163 0 : } else if (param.BatchSendRecvDataDes.curMode == BatchSendRecvCurMode::SEND) {
164 0 : CHK_RET(LocalNotify::Post(
165 : algResResp_->slaveStreams[STREAM_INDEX_0], dispatcher_, algResResp_->notifiesMain[STREAM_INDEX_0],
166 : PROF_STAGE_0));
167 0 : } else if (param.BatchSendRecvDataDes.curMode == BatchSendRecvCurMode::RECV) {
168 0 : CHK_RET(LocalNotify::Post(
169 : algResResp_->slaveStreams[STREAM_INDEX_1], dispatcher_, algResResp_->notifiesMain[STREAM_INDEX_1],
170 : PROF_STAGE_1));
171 : }
172 :
173 0 : CHK_RET(LaunchTaskExtend(dispatcher_, param.stream, algResResp_->slaveStreams));
174 0 : HCCL_INFO("[info][print] LaunchTaskExtend success.");
175 0 : HCCL_INFO(
176 : "tag[%s] BatchSendRecv Executor orchestrate success, take time [%lld]us.", param.tag.c_str(),
177 : DURATION_US(TIME_NOW() - startut));
178 0 : return HCCL_SUCCESS;
179 0 : }
180 :
181 0 : HcclResult CollBatchSendRecvRetryExecutor::RunLoop(
182 : OpParam& param, AlgResourceResponse& algRes, const std::vector<HcclSendRecvItem*>& sendRecvPair)
183 : {
184 : // 判断当前需执行的算子
185 0 : std::vector<HcclSendRecvItem*> curSendRecvPair;
186 0 : if (param.BatchSendRecvDataDes.curMode == BatchSendRecvCurMode::SEND) {
187 0 : curSendRecvPair.push_back(sendRecvPair[0]);
188 0 : } else if (param.BatchSendRecvDataDes.curMode == BatchSendRecvCurMode::RECV) {
189 0 : curSendRecvPair.push_back(sendRecvPair[sendRecvPair.size() - 1]);
190 : } else {
191 0 : curSendRecvPair = sendRecvPair;
192 : }
193 :
194 : // 执行当前需执行的算子
195 0 : for (const auto& itemPtr : curSendRecvPair) {
196 0 : if (static_cast<bool>(param.BatchSendRecvDataDes.isDirectRemoteRank[itemPtr->remoteRank])) {
197 : // device direct链路的任务会在host侧下发,此处需要跳过
198 0 : continue;
199 : }
200 0 : HCCL_INFO("[CollBatchSendRecvRetryExecutor][RunLoop] remoteRank %u", itemPtr->remoteRank);
201 0 : if (itemPtr->sendRecvType == HcclSendRecvType::HCCL_SEND) {
202 0 : CHK_RET(CalcSendSlices(algRes, itemPtr));
203 0 : } else if (itemPtr->sendRecvType == HcclSendRecvType::HCCL_RECV) {
204 0 : CHK_RET(CalcRecvSlices(algRes, itemPtr));
205 : } else {
206 0 : HCCL_ERROR("[CollBatchSendRecvRetryExecutor][RunLoop] sendRecvType is Wrong.");
207 0 : return HCCL_E_PARA;
208 : }
209 : }
210 :
211 0 : u32 loopInOnceLaunch = 0;
212 : // 每隔200个loop launch一次
213 0 : while (!sendDataSilces_.empty() || !recvDataSilces_.empty()) {
214 0 : if (!sendDataSilces_.empty()) {
215 0 : CHK_RET(ProcessSendDataSlice(algResResp_->slaveStreams[STREAM_INDEX_0], false, true));
216 0 : sendDataSilces_.pop_front();
217 : }
218 0 : if (!recvDataSilces_.empty()) {
219 0 : CHK_RET(ProcessRecvDataSlice(algResResp_->slaveStreams[STREAM_INDEX_1], true));
220 0 : recvDataSilces_.pop_front();
221 : }
222 0 : loopInOnceLaunch++;
223 0 : if (loopInOnceLaunch == MAX_LOOP_IN_ONCE_LAUNCH || (sendDataSilces_.empty() && recvDataSilces_.empty())) {
224 0 : CHK_RET(LaunchTaskExtend(dispatcher_, param.stream, algResResp_->slaveStreams));
225 0 : HCCL_INFO(
226 : "[BatchSendRecv] LaunchTaskExtend, unprocessed send slices[%u], recv slices[%u].",
227 : sendDataSilces_.size(), recvDataSilces_.size());
228 0 : loopInOnceLaunch = 0;
229 : }
230 : }
231 0 : return HCCL_SUCCESS;
232 0 : }
233 :
234 0 : HcclResult CollBatchSendRecvRetryExecutor::CalcStreamNum(u32& streamNum)
235 : {
236 0 : streamNum = LEVEL0_PLANE_NUM_IN_NPRING_DOUBLE;
237 0 : HCCL_INFO("[CollBatchSendRecvRetryExecutor][CalcScratchMemSize] tag_[%s], streamNum[%u].", tag_.c_str(), streamNum);
238 0 : return HCCL_SUCCESS;
239 : }
240 :
241 0 : HcclResult CollBatchSendRecvRetryExecutor::CalcSendSlices(AlgResourceResponse& algRes, HcclSendRecvItem* sendRecvItem)
242 : {
243 0 : HCCL_INFO(
244 : "[CollBatchSendRecvExecutor][CalcSendSlices] tag[%s], remoteRank[%u], buf[%p], count[%llu],"
245 : "dataType[%s], sendRecvType[%d].",
246 : tag_.c_str(), sendRecvItem->remoteRank, sendRecvItem->buf, sendRecvItem->count,
247 : GetDataTypeEnumStr(sendRecvItem->dataType).c_str(), sendRecvItem->sendRecvType);
248 0 : u8* curInputPtr = static_cast<u8*>(sendRecvItem->buf);
249 0 : CHK_PTR_NULL(curInputPtr);
250 0 : u32 unitSize = SIZE_TABLE[sendRecvItem->dataType];
251 0 : u64 maxCountPerLoop = CalcSendLoopMaxCount(const_cast<DeviceMem&>(algRes.cclInputMem), unitSize);
252 :
253 0 : for (u64 countLeft = sendRecvItem->count, curCount = 0, curOffset = 0; countLeft > 0; countLeft -= curCount) {
254 0 : curInputPtr += curOffset;
255 0 : curCount = (countLeft > maxCountPerLoop) ? maxCountPerLoop : countLeft;
256 0 : u64 curSize = curCount * unitSize; // 单位:字节
257 0 : sendDataSilces_.emplace_back(curInputPtr, curSize, sendRecvItem->remoteRank);
258 0 : HCCL_DEBUG(
259 : "[CollBatchSendRecvExecutor][CalcSendSlices] tag[%s], slice userAddr[%p], slice size[%llu].", tag_.c_str(),
260 : curInputPtr, curSize);
261 0 : curOffset = curSize;
262 : }
263 0 : return HCCL_SUCCESS;
264 : }
265 :
266 0 : HcclResult CollBatchSendRecvRetryExecutor::CalcRecvSlices(AlgResourceResponse& algRes, HcclSendRecvItem* sendRecvItem)
267 : {
268 0 : HCCL_INFO(
269 : "[CollBatchSendRecvRetryExecutor][CalcSendSlices] tag[%s], remoteRank[%u], buf[%p], count[%llu],"
270 : "dataType[%s], sendRecvType[%d].",
271 : tag_.c_str(), sendRecvItem->remoteRank, sendRecvItem->buf, sendRecvItem->count,
272 : GetDataTypeEnumStr(sendRecvItem->dataType).c_str(), sendRecvItem->sendRecvType);
273 0 : u8* curOutputPtr = static_cast<u8*>(sendRecvItem->buf);
274 0 : CHK_PTR_NULL(curOutputPtr);
275 0 : u32 unitSize = SIZE_TABLE[sendRecvItem->dataType];
276 0 : u64 maxCountPerLoop = CalcRecvLoopMaxCount(const_cast<DeviceMem&>(algRes.cclOutputMem), unitSize);
277 0 : HCCL_DEBUG("[CollBatchSendRecvRetryExecutor][CalcSendSlices]maxCountPerLoop is %llu", maxCountPerLoop);
278 :
279 0 : for (u64 countLeft = sendRecvItem->count, curCount = 0, curOffset = 0; countLeft > 0; countLeft -= curCount) {
280 0 : curOutputPtr += curOffset;
281 0 : curCount = (countLeft > maxCountPerLoop) ? maxCountPerLoop : countLeft;
282 0 : u64 curSize = curCount * unitSize; // 单位:字节
283 0 : recvDataSilces_.emplace_back(curOutputPtr, curSize, sendRecvItem->remoteRank);
284 0 : HCCL_DEBUG(
285 : "[CollBatchSendRecvRetryExecutor][CalcRecvSlices] tag[%s], slice userAddr[%p], slice size[%llu].",
286 : tag_.c_str(), curOutputPtr, curSize);
287 0 : curOffset = curSize;
288 : }
289 0 : return HCCL_SUCCESS;
290 : }
291 :
292 : REGISTER_EXEC("BatchSendRecvRetry", BatchSendRecvRetryExecutor, CollBatchSendRecvRetryExecutor);
293 : } // namespace hccl
|