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 "aicpu_allgather.h"
12 : #include "common/aicpu_hccl_common.h"
13 :
14 : bool AicpuAllgather::isMCFirstCall = true;
15 :
16 32 : HcclResult AicpuAllgather::RunAlgorithm(
17 : HcclReduceOp opType, void* sendBuffer, void* recvBuffer, u64 dataCount, HcclDataType dataType, u64 strideLen,
18 : AivAicpuOpParam* nextTask)
19 : {
20 32 : CHK_PTR_NULL(ctx_);
21 32 : switch (ctx_->commAlg) {
22 24 : case CommAlgType::COMM_ALG_FULL_MESH: {
23 24 : if (ctx_->commLen < ctx_->windowSize / AC_DEFAULT_WINDOW_DIM) {
24 16 : if (nextTask == nullptr
25 22 : || DataUnitSize(nextTask->hcclDataType) * nextTask->count
26 6 : > ctx_->windowSize / AC_DEFAULT_WINDOW_DIM) {
27 10 : return RunAllGathervMC(opType, sendBuffer, recvBuffer, dataCount, dataType, strideLen, nullptr);
28 : } else {
29 6 : return RunAllGathervMC(opType, sendBuffer, recvBuffer, dataCount, dataType, strideLen, nextTask);
30 : }
31 : }
32 8 : return RunAllGatherv(opType, sendBuffer, recvBuffer, dataCount, dataType, strideLen);
33 : }
34 8 : case CommAlgType::COMM_ALG_DOUBLE_RING: {
35 8 : return RunDoubleRingAllGather(
36 8 : opType, reinterpret_cast<u64>(sendBuffer), reinterpret_cast<u64>(recvBuffer), dataCount, dataType);
37 : }
38 0 : default: {
39 0 : HCCL_ERROR("CommAlg %d is not supported.", ctx_->commAlg);
40 0 : return HCCL_E_NOT_SUPPORT;
41 : }
42 : }
43 : }
44 :
45 8 : HcclResult AicpuAllgather::RunAllGatherv(
46 : HcclReduceOp opType, void* sendBuffer, void* recvBuffer, u64 dataCount, HcclDataType dataType, u64 strideCnt)
47 : {
48 8 : CHK_PTR_NULL(ctx_);
49 :
50 : // dataCount 为 gather前的数据量
51 8 : u64 windowSize = ctx_->windowSize;
52 8 : u32 unitSize = ctx_->unitSize;
53 8 : u64 maxCountPerLoop = windowSize / unitSize; // 中转内存单次最多能够接受的output count
54 :
55 8 : u8* curInputPtr = static_cast<u8*>(sendBuffer);
56 8 : u8* curOutputPtr = static_cast<u8*>(recvBuffer);
57 8 : u64 inputOffset = 0;
58 8 : u64 outputOffset = 0;
59 8 : u64 countLeft = dataCount;
60 :
61 8 : u64 displs[AC_MAX_RANK_NUM] = {0};
62 72 : for (u32 i = 0; i < ctx_->rankNum; i++) {
63 64 : displs[i] = i * strideCnt * ctx_->unitSize;
64 : }
65 :
66 8 : u32 loopIdx = 0;
67 16 : while (countLeft > 0) {
68 8 : curInputPtr += inputOffset;
69 8 : curOutputPtr += outputOffset;
70 8 : u64 curCount = (countLeft > maxCountPerLoop) ? maxCountPerLoop : countLeft;
71 8 : u64 curSize = curCount * unitSize; // 单位 byte
72 :
73 8 : HCCL_DEBUG(
74 : "RunAllGatherv: loop %u, curInputPtr[%p], curOutputPtr[%p], curCount[%llu], curSize[%llu]", loopIdx++,
75 : curInputPtr, curOutputPtr, curCount, curSize);
76 :
77 : // 1. 片内数据拷贝 snd->win
78 8 : CHK_RET(TaskOrchestrator::SelfCpySnd2Win(curInputPtr, curSize, 0, 0, HCCL_REDUCE_RESERVED, dataType));
79 : // 2. 前同步
80 8 : CHK_RET(TaskOrchestrator::DoPreSync());
81 :
82 : // 缺省情况allgather时op_type没用,融合时aic设置op_type为1时表示不需要本卡数据。
83 8 : if (opType != HCCL_REDUCE_PROD) {
84 0 : CHK_RET(TaskOrchestrator::SelfCpyWin2Rcv(
85 : curOutputPtr, curSize, 0, displs[ctx_->rankId], HCCL_REDUCE_RESERVED, dataType));
86 : }
87 :
88 : // 3. 跨片SDMA,其他win->当前rcv buff
89 8 : CHK_RET(
90 : TaskOrchestrator::IpcCpyWin2Rcv(curOutputPtr, curSize, nullptr, displs, HCCL_REDUCE_RESERVED, dataType));
91 :
92 : // 4. 后同步
93 8 : CHK_RET(TaskOrchestrator::DoPostSync());
94 :
95 8 : CHK_RET(TaskOrchestrator::LaunchTasks());
96 :
97 8 : countLeft -= curCount;
98 8 : inputOffset = curSize;
99 8 : outputOffset = curSize;
100 : }
101 8 : return HCCL_SUCCESS;
102 : }
103 :
104 22 : u64 AicpuAllgather::GetWindowOffset(u32 curTurnCnt, u64 curSize, u64 strideCnt, u64 recvBuffer)
105 : {
106 : (void)recvBuffer;
107 22 : u64 windowOffset = 0;
108 22 : u64 bufferFlag = static_cast<uint64_t>(curTurnCnt % AC_DEFAULT_WINDOW_DIM);
109 22 : u64 recvOffset
110 22 : = (static_cast<uint64_t>(ctx_->rankId) * strideCnt * static_cast<uint64_t>(ctx_->unitSize)) % HCCL_COPY_ALIGN;
111 22 : u64 windowOffsetTmp
112 22 : = bufferFlag * (ctx_->windowSize / AC_DEFAULT_WINDOW_DIM / HCCL_COPY_ALIGN + 1) * HCCL_COPY_ALIGN + recvOffset;
113 22 : if ((bufferFlag == 0 && (recvOffset + curSize) < ctx_->windowSize / AC_DEFAULT_WINDOW_DIM)
114 13 : || (bufferFlag == 1 && (windowOffsetTmp + curSize) < ctx_->windowSize)) {
115 22 : windowOffset = windowOffsetTmp;
116 : } else {
117 0 : windowOffset = bufferFlag * (ctx_->windowSize / AC_DEFAULT_WINDOW_DIM);
118 : }
119 22 : return windowOffset;
120 : }
121 :
122 16 : HcclResult AicpuAllgather::RunAllGathervMC(
123 : HcclReduceOp opType, void* sendBuffer, void* recvBuffer, u64 dataCount, HcclDataType dataType, u64 strideCnt,
124 : AivAicpuOpParam* nextTask)
125 : {
126 16 : u64 displs[AC_MAX_RANK_NUM] = {0};
127 144 : for (u32 i = 0; i < ctx_->rankNum; i++) {
128 128 : displs[i] = i * strideCnt * ctx_->unitSize;
129 : }
130 :
131 16 : u64 curSize = ctx_->unitSize * dataCount;
132 16 : u64 windowOffset = 0;
133 16 : if (curSize < ctx_->windowSize / AC_DEFAULT_WINDOW_DIM) {
134 16 : windowOffset = GetWindowOffset(ctx_->curTurnCnt, curSize, strideCnt, reinterpret_cast<uint64_t>(recvBuffer));
135 : }
136 16 : HCCL_INFO(
137 : "current task: windowSize[%llu], windowOffset[%llu], dataCount[%llu], "
138 : "curSize[%llu], strideCnt[%llu], gatherOut [%#llx]",
139 : ctx_->windowSize, windowOffset, dataCount, curSize, strideCnt, ctx_->gatherOut);
140 :
141 : // 1. 片内数据拷贝 snd->win
142 16 : if (isMCFirstCall) {
143 10 : CHK_RET(TaskOrchestrator::SelfCpySnd2Win(sendBuffer, curSize, 0, windowOffset, HCCL_REDUCE_RESERVED, dataType));
144 10 : isMCFirstCall = false;
145 : }
146 :
147 : // 2. 前同步
148 16 : CHK_RET(TaskOrchestrator::DoPreSync());
149 :
150 : // 缺省情况allgather时op_type没用,融合时aic设置op_type为1时表示不需要本卡数据。
151 16 : if (opType != HCCL_REDUCE_PROD) {
152 0 : CHK_RET(TaskOrchestrator::SelfCpyWin2Rcv(
153 : recvBuffer, curSize, windowOffset, displs[ctx_->rankId], HCCL_REDUCE_RESERVED, dataType));
154 : }
155 :
156 : // 3. 跨片SDMA,其他win->当前rcv buff
157 16 : u64 winOffsets[AC_MAX_RANK_NUM] = {0};
158 144 : for (u32 i = 0; i < ctx_->rankNum; i++) {
159 128 : winOffsets[i] = (ctx_->curTurnCnt % AC_DEFAULT_WINDOW_DIM)
160 128 : * (ctx_->windowSize / AC_DEFAULT_WINDOW_DIM / HCCL_COPY_ALIGN + 1) * HCCL_COPY_ALIGN
161 128 : + (i * strideCnt * ctx_->unitSize) % HCCL_COPY_ALIGN;
162 : }
163 16 : CHK_RET(TaskOrchestrator::IpcCpyWin2Rcv(recvBuffer, curSize, winOffsets, displs, HCCL_REDUCE_RESERVED, dataType));
164 :
165 16 : if (nextTask != nullptr) {
166 : // 1. 片内数据拷贝 snd->win
167 6 : u64 nextSize = DataUnitSize(nextTask->hcclDataType) * nextTask->count;
168 6 : u64 nextWindowOffset = 0;
169 6 : if (curSize < ctx_->windowSize / AC_DEFAULT_WINDOW_DIM) {
170 6 : nextWindowOffset = GetWindowOffset(ctx_->curTurnCnt + 1, nextSize, strideCnt, nextTask->recvBuffer);
171 : }
172 :
173 6 : HCCL_INFO(
174 : "curTurnCnt[%u], windowSize[%llu], windowOffset[%llu], dataCount[%llu], nextSize[%llu]", ctx_->curTurnCnt,
175 : ctx_->windowSize, nextWindowOffset, nextTask->count, nextSize);
176 :
177 6 : CHK_RET(TaskOrchestrator::SelfCpySnd2Win(
178 : reinterpret_cast<void*>(nextTask->sendBuffer), nextSize, 0, nextWindowOffset, HCCL_REDUCE_RESERVED,
179 : nextTask->hcclDataType));
180 : } else {
181 10 : isMCFirstCall = true;
182 : }
183 :
184 : // 4. 后同步
185 16 : CHK_RET(TaskOrchestrator::DoPostSync());
186 :
187 16 : CHK_RET(TaskOrchestrator::LaunchTasks());
188 :
189 16 : return HCCL_SUCCESS;
190 : }
191 :
192 112 : HcclResult AicpuAllgather::GenRingTask(
193 : [[maybe_unused]] HcclReduceOp opType, u64 sndAddr, u64 rcvAddr, u64 gatherSize, HcclDataType dataType,
194 : uint32_t streamId, bool isClockwise, uint32_t step, bool isWindowLast) const
195 : {
196 112 : const u64 winIn = ctx_->rankInfo[rankId_].window + (isClockwise ? 0U : ctx_->windowSize / RING_NUM);
197 112 : const u64 winOut = ctx_->rankInfo[rankId_].windowOut + (isClockwise ? 0U : ctx_->windowSize / RING_NUM);
198 112 : const uint32_t preRankId = isClockwise ? (rankId_ + rankNum_ - 1U) % rankNum_ : (rankId_ + 1U) % rankNum_;
199 112 : const uint32_t postRankId = isClockwise ? (rankId_ + 1U) % rankNum_ : (rankId_ + rankNum_ - 1U) % rankNum_;
200 112 : const u64 preWinIn = ctx_->rankInfo[preRankId].window + (isClockwise ? 0U : ctx_->windowSize / RING_NUM);
201 112 : const u64 preWinOut = ctx_->rankInfo[preRankId].windowOut + (isClockwise ? 0U : ctx_->windowSize / RING_NUM);
202 112 : const bool evenStep = (step % 2U == 0U); // 2: 环上winIn winOut轮流接收, 奇数轮in->out, 偶数轮out->in
203 :
204 112 : HcclResult ret = HCCL_SUCCESS;
205 112 : if (step == 1U) { // 首轮send->winIn
206 16 : ret = AicpuDispatcher::CopyData(streamId, sndAddr, winIn, gatherSize, dataType, HCCL_REDUCE_RESERVED, rankId_);
207 16 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("Turn %u step %u send to clock winIn failed", turn_, step), ret);
208 : }
209 : // 片间同步 notify后卡 wait前卡
210 112 : ret = AicpuDispatcher::SignalRecord(streamId, postRankId, AicpuDispatcher::IPC, AicpuDispatcher::PRE_SYNC);
211 112 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("Turn %u step %u notify post rank failed", turn_, step), ret);
212 112 : ret = AicpuDispatcher::SignalWait(streamId, preRankId, AicpuDispatcher::IPC, AicpuDispatcher::PRE_SYNC);
213 112 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("Turn %u step %u wait pre rank failed", turn_, step), ret);
214 :
215 : // 片间memcpy 前卡window->recv
216 112 : auto recvBufAddr = rcvAddr
217 112 : + (isClockwise ? (rankId_ + rankNum_ - step) % rankNum_ : (rankId_ + step) % rankNum_)
218 112 : * ctx_->totalCnt * unitSize_;
219 112 : ret = AicpuDispatcher::CopyData(
220 : streamId, evenStep ? preWinOut : preWinIn, recvBufAddr, gatherSize, dataType, HCCL_REDUCE_RESERVED, preRankId);
221 112 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("Turn %u step %u cpy pre window failed", turn_, step), ret);
222 :
223 112 : if (isClockwise) {
224 56 : if (!((step == rankNum_ - 1U) && isWindowLast)) {
225 48 : CHK_RET(AicpuDispatcher::SignalWait(
226 : streamId, (rankId_ + 1U) % rankNum_, AicpuDispatcher::NO_IPC, AicpuDispatcher::POST_SYNC));
227 : }
228 : // ccore notify
229 56 : if ((step < rankNum_ - 1U) && isWindowLast) {
230 48 : ret = AicpuDispatcher::AddCcoreNotify(streamId, turn_ * (rankNum_ - 1U) + step);
231 48 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("Turn %u step %u add ccore notify failed", turn_, step), ret);
232 : }
233 : } else {
234 56 : if (!((step == rankNum_ - 1U) && isWindowLast)) {
235 48 : CHK_RET(AicpuDispatcher::SignalRecord(
236 : streamId, (rankId_ + 1U) % rankNum_, AicpuDispatcher::NO_IPC, AicpuDispatcher::POST_SYNC));
237 : }
238 : }
239 : // recv->windowout
240 112 : if (step < rankNum_ - 1U) {
241 96 : ret = AicpuDispatcher::CopyData(
242 96 : streamId, recvBufAddr, evenStep ? winIn : winOut, gatherSize, dataType, HCCL_REDUCE_RESERVED, rankId_);
243 96 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("Turn %u step %u cpy pre window failed", turn_, step), ret);
244 : }
245 :
246 : // 片间同步 notify前卡 wait后卡
247 112 : ret = AicpuDispatcher::SignalRecord(streamId, preRankId, AicpuDispatcher::IPC, AicpuDispatcher::POST_SYNC);
248 112 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("Turn %u step %u notify pre rank failed", turn_, step), ret);
249 112 : ret = AicpuDispatcher::SignalWait(streamId, postRankId, AicpuDispatcher::IPC, AicpuDispatcher::POST_SYNC);
250 112 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("Turn %u step %u wait post rank failed", turn_, step), ret);
251 112 : return HCCL_SUCCESS;
252 : }
253 :
254 8 : HcclResult AicpuAllgather::RunDoubleRingAllGather(
255 : HcclReduceOp opType, u64 sendBuffer, u64 recvBuffer, u64 dataCount, HcclDataType dataType) const
256 : {
257 8 : const u64 gatherSize = dataCount * unitSize_;
258 8 : if (gatherSize > ctx_->windowSize / RING_NUM) {
259 0 : HCCL_INFO("Tile gather size %lu max less than window size %lu/RING_NUM", gatherSize, ctx_->windowSize);
260 : }
261 8 : const u64 rankDataSize = ctx_->totalCnt * unitSize_;
262 : thread_local static u64 sndAddr[RING_NUM] = {0UL};
263 : thread_local static u64 rcvAddr[RING_NUM] = {0UL};
264 8 : if (turn_ == 0U) { // 首个tile时初始化snd rcv地址
265 1 : sndAddr[0] = sendBuffer;
266 1 : sndAddr[1] = sendBuffer + rankDataSize;
267 1 : rcvAddr[0] = recvBuffer;
268 1 : rcvAddr[1] = recvBuffer + rankDataSize;
269 : }
270 :
271 8 : HCCL_INFO(
272 : "DR AllGather snd addr %p %p, rcv addr %p %p, size %lu", sndAddr[0], sndAddr[1], rcvAddr[0], rcvAddr[1],
273 : gatherSize);
274 8 : uint32_t mainStream = rankId_;
275 8 : uint32_t subStream = (rankId_ + 1U) % rankNum_;
276 8 : HcclResult ret = HCCL_SUCCESS;
277 :
278 8 : u64 maxCountPerLoop = ctx_->windowSize / RING_NUM / unitSize_; // 中转内存单次最多能够接受的output count
279 8 : u64 countLeft = dataCount;
280 8 : bool isWindowFirst = true;
281 8 : bool isWindowLast = false;
282 : // windowSize循环
283 8 : u32 loopIdx = 0;
284 16 : while (countLeft > 0) {
285 8 : u64 curCount = (countLeft > maxCountPerLoop) ? maxCountPerLoop : countLeft;
286 8 : u64 curSize = curCount * unitSize_; // 单位 byte
287 :
288 8 : HCCL_DEBUG(
289 : "DR AllGather: loop %u, snd addr[%u %u], rcv addr[%u %u], curCount[%llu], curSize[%llu]", loopIdx++,
290 : sndAddr[0], sndAddr[1], rcvAddr[0], rcvAddr[1], curCount, curSize);
291 :
292 8 : sndAddr[1] -= curSize;
293 8 : rcvAddr[1] -= curSize; // 逆时针输出在执行前向上偏移
294 :
295 8 : isWindowLast = ((countLeft - curCount) == 0);
296 :
297 : // 顺、逆时针各平移ranknum-1次
298 64 : for (uint32_t step = 1U; step <= rankNum_ - 1U; step++) {
299 56 : if (isWindowFirst) {
300 : // ccore wait
301 56 : u64 waitAddr = ctx_->workSpaceAddr + ctx_->notifyOff + offsetof(AivAicpuOpParam, sendCnt);
302 112 : ret = AicpuDispatcher::AddCcoreWait(
303 56 : mainStream, waitAddr, turn_ * (rankNum_ - 1U) + step,
304 56 : (turn_ + 1U >= ctx_->totalTurnCnt) && (step == rankNum_ - 1U));
305 56 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("Turn %u step %u add ccore wait failed", turn_, step), ret);
306 : }
307 :
308 : // 主->从
309 56 : CHK_RET(AicpuDispatcher::SignalRecord(
310 : mainStream, subStream, AicpuDispatcher::NO_IPC, AicpuDispatcher::PRE_SYNC));
311 56 : CHK_RET(
312 : AicpuDispatcher::SignalWait(subStream, subStream, AicpuDispatcher::NO_IPC, AicpuDispatcher::PRE_SYNC));
313 :
314 : // 顺时针环
315 56 : CHK_RET(
316 : GenRingTask(opType, sndAddr[0], rcvAddr[0], curSize, dataType, mainStream, true, step, isWindowLast));
317 :
318 : // 逆时针环
319 56 : CHK_RET(
320 : GenRingTask(opType, sndAddr[1], rcvAddr[1], curSize, dataType, subStream, false, step, isWindowLast));
321 :
322 : // 缺省情况allgather时op_type没用,融合时aic设置op_type为1时表示不需要本卡数据。
323 56 : if ((step == rankNum_ - 1U) && isWindowLast) { // 拷贝全量本卡直接从send->recv
324 8 : if (opType != HCCL_REDUCE_PROD) {
325 0 : ret = AicpuDispatcher::CopyData(
326 0 : mainStream, sendBuffer, recvBuffer + rankId_ * rankDataSize, rankDataSize, dataType,
327 0 : HCCL_REDUCE_RESERVED, rankId_);
328 0 : CHK_PRT_RET(
329 : ret != HCCL_SUCCESS, HCCL_ERROR("Turn %u step %u send to recv winIn failed", turn_, step), ret);
330 : }
331 :
332 : // 从->主
333 8 : CHK_RET(AicpuDispatcher::SignalRecord(
334 : subStream, subStream, AicpuDispatcher::NO_IPC, AicpuDispatcher::POST_SYNC));
335 8 : CHK_RET(AicpuDispatcher::SignalWait(
336 : mainStream, subStream, AicpuDispatcher::NO_IPC, AicpuDispatcher::POST_SYNC));
337 :
338 8 : ret = TaskOrchestrator::AddBarrier(mainStream, rankId_, rankNum_);
339 8 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("Turn %u step %u add barrier failed.", turn_, step), ret);
340 :
341 : // ccore notify 非step7由GenRingTask处理AddCcoreNotify
342 8 : ret = AicpuDispatcher::AddCcoreNotify(mainStream, turn_ * (rankNum_ - 1U) + step);
343 8 : CHK_PRT_RET(
344 : ret != HCCL_SUCCESS, HCCL_ERROR("Turn %u step %u add ccore notify failed", turn_, step), ret);
345 : }
346 : }
347 8 : isWindowFirst = false;
348 8 : countLeft -= curCount;
349 8 : sndAddr[0] += curSize;
350 8 : rcvAddr[0] += curSize; // 顺时针输出在执行后向下偏移
351 : }
352 :
353 8 : ret = TaskOrchestrator::LaunchTasks();
354 8 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("Launch tasks failed"), ret);
355 8 : return HCCL_SUCCESS;
356 : }
|