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_dmy_cal_allreduce.h"
12 : #include "common/aicpu_kfc_def.h"
13 :
14 : namespace {
15 : constexpr u64 ALL_REDUCE_THRESHOLD = 1024 * 1024; // allreduce确定性计算,小于该值选择AL算法,否则选择RLA算法
16 : constexpr u64 ALL_REDUCE_THRESHOLD_UT = 1024; // allreduce确定性计算ut测试值
17 : } // namespace
18 :
19 1 : HcclResult AicpuDmyCalAllreduce::RunAlgorithm(
20 : HcclReduceOp opType, void* sendBuffer, void* recvBuffer, u64 dataCount, HcclDataType dataType,
21 : [[maybe_unused]] u64 strideLen, AivAicpuOpParam* /* nextTask */)
22 : {
23 1 : CHK_PTR_NULL(ctx_);
24 1 : HcclResult ret = HCCL_SUCCESS;
25 : #ifdef RUN_TEST
26 : ctx_->windowSize = ALL_REDUCE_THRESHOLD_UT * ctx_->unitSize;
27 : if (ctx_->commLen < ALL_REDUCE_THRESHOLD_UT) {
28 : #else
29 1 : if (ctx_->commLen < ALL_REDUCE_THRESHOLD) {
30 : #endif
31 0 : ret = RunAllReduceAL(opType, sendBuffer, recvBuffer, dataCount, dataType);
32 : } else {
33 1 : const u64 tailCount = dataCount % (HCCL_COPY_ALIGN / ctx_->unitSize);
34 1 : const u64 tileCount = dataCount - tailCount;
35 1 : CHK_RET(RunAllReduceRLA(opType, sendBuffer, recvBuffer, tileCount, dataType));
36 1 : const u64 offset = tileCount * ctx_->unitSize;
37 1 : CHK_RET(RunAllReduceAL(
38 : opType, static_cast<u8*>(sendBuffer) + offset, static_cast<u8*>(recvBuffer) + offset, tailCount, dataType));
39 : }
40 1 : return ret;
41 : }
42 :
43 1 : HcclResult AicpuDmyCalAllreduce::RunAllReduceAL(
44 : HcclReduceOp opType, void* sendBuffer, void* recvBuffer, u64 dataCount, HcclDataType dataType)
45 : {
46 1 : if (dataCount == 0UL) {
47 1 : return HCCL_SUCCESS;
48 : }
49 0 : u32 unitSize = ctx_->unitSize;
50 0 : u8* curInputPtr = static_cast<u8*>(sendBuffer);
51 0 : u8* curOutputPtr = static_cast<u8*>(recvBuffer);
52 :
53 0 : u64 windowSlices[AC_MAX_RANK_NUM] = {0};
54 0 : u64 curCount = dataCount;
55 0 : u64 curSize = curCount * unitSize;
56 :
57 0 : HCCL_INFO(
58 : "RunDmyCalAllreduceAL:curInputPtr[%p], curOutputPtr[%p], curCount[%llu], curSize[%llu]", curInputPtr,
59 : curOutputPtr, curCount, curSize);
60 :
61 0 : for (u32 i = 0; i < ctx_->rankNum; i++) {
62 0 : windowSlices[i] = i * curSize;
63 : }
64 :
65 : // 1. 片内数据拷贝 snd->win
66 0 : TaskOrchestrator::SelfCpySnd2Win(
67 0 : curInputPtr, curSize, 0, windowSlices[ctx_->rankId], HCCL_REDUCE_RESERVED, dataType);
68 :
69 : // 2. 前同步
70 0 : TaskOrchestrator::DoPreSync();
71 :
72 : // 3. 跨片SDMA,send->对端win
73 0 : TaskOrchestrator::IpcCpySnd2Win(
74 0 : curInputPtr, curSize, static_cast<u64>(0), windowSlices[ctx_->rankId], HCCL_REDUCE_RESERVED, dataType);
75 :
76 : // 4. 后同步
77 0 : TaskOrchestrator::DoPostSync();
78 :
79 : // 5. 折半计算
80 0 : TaskOrchestrator::SelfLocalReduce(curSize, opType, dataType);
81 :
82 : // 6. 前同步
83 0 : TaskOrchestrator::DoPreSync();
84 :
85 : // 7. 片内数据拷贝 本端win->当前rcv buff
86 0 : TaskOrchestrator::SelfCpyWin2Rcv(curOutputPtr, curSize, 0, 0, HCCL_REDUCE_RESERVED, dataType);
87 :
88 : // 8. 后同步
89 0 : TaskOrchestrator::DoPostSync();
90 :
91 0 : TaskOrchestrator::LaunchTasks();
92 0 : return HCCL_SUCCESS;
93 : }
94 :
95 1 : HcclResult AicpuDmyCalAllreduce::RunAllReduceRLA(
96 : HcclReduceOp opType, void* sendBuffer, void* recvBuffer, u64 dataCount, HcclDataType dataType) const
97 : {
98 1 : u64 windowSize = ctx_->windowSize;
99 1 : u32 unitSize = ctx_->unitSize;
100 1 : u64 maxCountPerLoop = windowSize / unitSize;
101 :
102 1 : u8* curInputPtr = static_cast<u8*>(sendBuffer);
103 1 : u8* curOutputPtr = static_cast<u8*>(recvBuffer);
104 1 : u64 inputOffset = 0;
105 1 : u64 outputOffset = 0;
106 1 : u64 countLeft = dataCount;
107 1 : u64 windowSlices[AC_MAX_RANK_NUM] = {0};
108 :
109 1 : u32 loopIdx = 0;
110 1 : while (countLeft > 0) {
111 0 : curInputPtr += inputOffset;
112 0 : curOutputPtr += outputOffset;
113 0 : u64 curCount = (countLeft > maxCountPerLoop) ? maxCountPerLoop : countLeft;
114 0 : u64 curSize = curCount * unitSize;
115 0 : u64 sliceSize = curSize / ctx_->rankNum;
116 :
117 0 : HCCL_INFO(
118 : "DmyCalAllreduceRLA:loop %u, curInputPtr[%p], curOutputPtr[%p], curCount[%llu], curSize[%llu]", loopIdx++,
119 : curInputPtr, curOutputPtr, curCount, curSize);
120 :
121 0 : for (u32 i = 0; i < ctx_->rankNum; i++) {
122 0 : windowSlices[i] = i * sliceSize;
123 : }
124 :
125 : // 1. 片内数据拷贝 snd->win
126 0 : TaskOrchestrator::SelfCpySnd2Win(
127 0 : curInputPtr, sliceSize, windowSlices[ctx_->rankId], windowSlices[ctx_->rankId], HCCL_REDUCE_RESERVED,
128 : dataType);
129 : // 2. 前同步
130 0 : TaskOrchestrator::DoPreSync();
131 :
132 : // 3. 跨片SDMA,send->其他window
133 0 : TaskOrchestrator::IpcCpySnd2Win(
134 0 : curInputPtr, sliceSize, windowSlices, windowSlices[ctx_->rankId], HCCL_REDUCE_RESERVED, dataType);
135 :
136 : // 4. 后同步
137 0 : TaskOrchestrator::DoPostSync();
138 :
139 : // 5. 折半计算
140 0 : TaskOrchestrator::SelfLocalReduce(sliceSize, opType, dataType);
141 :
142 : // 6. 片内数据拷贝 本端win->当前rcv buff
143 0 : TaskOrchestrator::SelfCpyWin2Rcv(
144 0 : curOutputPtr, sliceSize, 0, windowSlices[ctx_->rankId], HCCL_REDUCE_RESERVED, dataType);
145 :
146 : // 7. 前同步
147 0 : TaskOrchestrator::DoPreSync();
148 :
149 : // 8. 跨片SDMA send->其他window
150 0 : TaskOrchestrator::IpcCpyWin2Rcv(curOutputPtr, sliceSize, nullptr, windowSlices, HCCL_REDUCE_RESERVED, dataType);
151 :
152 : // 9. 后同步
153 0 : TaskOrchestrator::DoPostSync();
154 :
155 0 : TaskOrchestrator::LaunchTasks();
156 :
157 0 : countLeft -= curCount;
158 0 : inputOffset = curSize;
159 0 : outputOffset = curSize;
160 : }
161 1 : return HCCL_SUCCESS;
162 : }
|