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 <algorithm>
12 : #include <mutex>
13 : #include <thread>
14 : #include <vector>
15 :
16 : #include "hccl_group.h"
17 :
18 : using namespace hccl;
19 :
20 : thread_local s32 hcclGroupDepth = 0;
21 : thread_local std::deque<std::shared_ptr<struct hcclAsyncJob>> hcclInitJobs;
22 : thread_local std::vector<HcclComm> hcclGroupCommList;
23 :
24 2 : HcclResult HcclLegacyGroupStart()
25 : {
26 2 : hcclGroupDepth++;
27 2 : HCCL_INFO("[HcclGroupStart] hcclGroupDepth=[%d]", hcclGroupDepth);
28 2 : return HCCL_SUCCESS;
29 : }
30 :
31 : namespace hccl {
32 :
33 0 : HcclResult initGroupPlanner(HcclComm comm)
34 : {
35 0 : hccl::hcclComm* hcclComm = static_cast<hccl::hcclComm*>(comm);
36 0 : std::shared_ptr<struct hcclKernelPlanner> planner = hcclComm->planner;
37 0 : u32 rankSize = INVALID_VALUE_RANKSIZE;
38 0 : hcclComm->GetRankSize(rankSize);
39 0 : planner->rankSize = rankSize;
40 0 : HCCL_DEBUG("[initGroupPlanner] ranksize: %d", rankSize);
41 :
42 0 : planner->nTasksP2p = 0;
43 0 : planner->nTasksColl = 0;
44 0 : return HCCL_SUCCESS;
45 0 : }
46 :
47 3 : HcclResult taskAppend(HcclComm comm, hcclOpInfo& info)
48 : {
49 3 : hccl::hcclComm* hcclComm = static_cast<hccl::hcclComm*>(comm);
50 3 : std::shared_ptr<struct hcclKernelPlanner> planner = hcclComm->planner;
51 3 : if (planner->nTasksP2p == -1) {
52 0 : initGroupPlanner(comm);
53 : }
54 :
55 3 : HcclResult ret = HCCL_SUCCESS;
56 3 : if (info.coll == HcclCMDType::HCCL_CMD_SEND || info.coll == HcclCMDType::HCCL_CMD_RECEIVE) {
57 2 : hcclComm->SetGroupMode(true);
58 2 : bool isSendOp = (info.coll == HcclCMDType::HCCL_CMD_SEND);
59 : HcclSendRecvItem item;
60 2 : item.sendRecvType = isSendOp ? HcclSendRecvType::HCCL_SEND : HcclSendRecvType::HCCL_RECV;
61 2 : item.buf = isSendOp ? info.sendbuff : info.recvbuff;
62 2 : item.count = isSendOp ? info.sendCount : info.recvCount;
63 2 : item.dataType = isSendOp ? info.sendType : info.recvType;
64 2 : item.remoteRank = info.root;
65 :
66 2 : planner->sendRecvInfo.push_back(item);
67 :
68 2 : if (planner->sendRecvMainStream == nullptr) { // 用第一条用户流作为主流
69 2 : planner->sendRecvMainStream = info.stream;
70 2 : HCCL_INFO("[TaskAppend] planner->sendRecvMainStream[%p]", planner->sendRecvMainStream);
71 : }
72 :
73 2 : planner->nTasksP2p += 1;
74 2 : } else {
75 1 : hcclOpInfo task = info;
76 1 : planner->collTaskQueue.push_back(task);
77 1 : planner->nTasksColl += 1;
78 : /*记录stream到planner*/
79 1 : planner->collStreams.insert(info.stream);
80 : }
81 :
82 3 : auto itComm = std::find(hcclGroupCommList.begin(), hcclGroupCommList.end(), comm);
83 3 : if (itComm == hcclGroupCommList.end()) {
84 3 : hcclGroupCommList.push_back(comm);
85 : }
86 3 : return ret;
87 3 : }
88 :
89 : HcclResult
90 2 : commInitTaskAppend(std::shared_ptr<struct hcclAsyncJob> job, HcclResult (*func)(struct hcclAsyncJob*), HcclComm* comm)
91 : {
92 2 : HCCL_INFO("[hcclAsyncJobEnqueue] add item to queue");
93 2 : CHK_PRT_RET(!job, HCCL_ERROR("[commInitTaskAppend] job is nullptr"), HCCL_E_INTERNAL);
94 : /*hcclAsyncLaunch只是将job放入队列,并不等待执行完成。groupLaunch->asyncJobLaunch中给每个job起一个线程去执行*/
95 1 : job->func = func;
96 1 : job->comm = comm;
97 1 : job->state = hcclGroupJobRunning;
98 1 : hcclInitJobs.push_back(job);
99 1 : return HCCL_SUCCESS;
100 : }
101 : } // namespace hccl
102 :
103 1 : void* hcclAsyncJobMain(void* arg)
104 : {
105 1 : struct hcclAsyncJob* job = (struct hcclAsyncJob*)arg;
106 1 : job->result = job->func(job); /*func是上层asyncjob里面设置的函数*/
107 1 : if (job->result == HCCL_SUCCESS) {
108 1 : HCCL_INFO("Function launch success");
109 : }
110 : /*加锁修改job->state为hcclGroupJobDone*/
111 1 : std::unique_lock<std::mutex> lock(job->mtx);
112 1 : job->state = hcclGroupJobDone;
113 1 : return arg;
114 1 : }
115 :
116 1 : static HcclResult asyncJobLaunch()
117 : {
118 1 : HCCL_DEBUG("[asyncJobLaunch] entered");
119 1 : HcclResult ret = HCCL_SUCCESS;
120 1 : bool jobsDone = false;
121 :
122 1 : if (!hcclInitJobs.empty()) {
123 2 : for (auto job : hcclInitJobs) {
124 1 : CHK_PRT_RET(!job, HCCL_ERROR("[asyncJobLaunch] job is nullptr"), HCCL_E_INTERNAL);
125 1 : job->thread.reset(new (std::nothrow) std::thread(&hcclAsyncJobMain, job.get()));
126 1 : CHK_PRT_RET(!job->thread, HCCL_ERROR("[asyncJobLaunch]threads reset failed "), HCCL_E_INTERNAL);
127 1 : }
128 :
129 : do { /*主线程轮询阻塞,等待所有线程上的asyncJob执行完成*/
130 1 : jobsDone = true;
131 2 : for (auto job : hcclInitJobs) {
132 : /*上面job执行线程可能并发修改state,在主线程里面要通过加线程锁来读取*/
133 1 : hcclGroupJobState_t state = hcclGroupJobJoined;
134 1 : std::unique_lock<std::mutex> lock(job->mtx);
135 1 : state = job->state;
136 :
137 1 : if (state == hcclGroupJobRunning) {
138 0 : jobsDone = false;
139 1 : } else if (state == hcclGroupJobDone) {
140 1 : job->thread->join();
141 1 : job->state = hcclGroupJobJoined;
142 1 : if (job->result != HCCL_SUCCESS && ret == HCCL_SUCCESS) {
143 0 : ret = job->result;
144 : }
145 : } else {
146 : /* safety check */
147 0 : CHK_PRT_RET(
148 : state != hcclGroupJobJoined, HCCL_ERROR("[asyncJobLaunch] state != hcclGroupJobJoined"),
149 : HCCL_E_INTERNAL);
150 : }
151 1 : }
152 : // Let preconnect threads progress.
153 1 : if (jobsDone == false)
154 0 : usleep(1);
155 1 : } while (jobsDone == false);
156 1 : hcclInitJobs.clear();
157 1 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[asyncJobLaunch] fail!"), ret);
158 : }
159 1 : return HCCL_SUCCESS;
160 : }
161 :
162 : // 计算集合通信算子的数据量(字节),用于按数据量降序排列(仅非V类算子参与排序)
163 1034 : u64 calcOpDataVolume(const hcclOpInfo& info)
164 : {
165 1034 : switch (info.coll) {
166 1012 : case HcclCMDType::HCCL_CMD_ALLGATHER:
167 : case HcclCMDType::HCCL_CMD_ALLTOALL:
168 : case HcclCMDType::HCCL_CMD_ALLREDUCE:
169 : case HcclCMDType::HCCL_CMD_BROADCAST:
170 : case HcclCMDType::HCCL_CMD_REDUCE:
171 1012 : return info.sendCount * SIZE_TABLE[info.sendType];
172 15 : case HcclCMDType::HCCL_CMD_REDUCE_SCATTER:
173 : case HcclCMDType::HCCL_CMD_SCATTER:
174 15 : return info.recvCount * SIZE_TABLE[info.recvType];
175 7 : default:
176 7 : return 0;
177 : }
178 : }
179 :
180 : // 对group任务排序:非V类按数据量降序在前,V类保持原序在后
181 11 : std::vector<hcclOpInfo> sortGroupTasks(const std::deque<hcclOpInfo>& tasks)
182 : {
183 11 : std::vector<hcclOpInfo> sorted(tasks.begin(), tasks.end());
184 129 : auto isVariant = [](const hcclOpInfo& op) {
185 124 : return op.coll == HcclCMDType::HCCL_CMD_ALLTOALLV || op.coll == HcclCMDType::HCCL_CMD_ALLTOALLVC
186 253 : || op.coll == HcclCMDType::HCCL_CMD_ALLGATHER_V || op.coll == HcclCMDType::HCCL_CMD_REDUCE_SCATTER_V;
187 : };
188 11 : auto itV = std::stable_partition(sorted.begin(), sorted.end(), [&isVariant](const hcclOpInfo& op) {
189 129 : return !isVariant(op);
190 : });
191 11 : std::stable_sort(sorted.begin(), itV, [](const hcclOpInfo& a, const hcclOpInfo& b) {
192 398 : return calcOpDataVolume(a) > calcOpDataVolume(b);
193 : });
194 11 : return sorted;
195 0 : }
196 :
197 1 : static HcclResult doLaunches(HcclComm comm)
198 : {
199 1 : hccl::hcclComm* hcclComm = static_cast<hccl::hcclComm*>(comm);
200 1 : std::shared_ptr<struct hcclKernelPlanner> planner = hcclComm->planner;
201 1 : HcclUs startutime = TIME_NOW();
202 1 : if (planner->nTasksP2p != 0) {
203 : // 将所有send/recv的任务打包作为一个集合通信算子来执行
204 1 : HCCL_INFO("HcclBatchSendRecvGroup, sendRecvInfo.size()[%u]", static_cast<u32>(planner->sendRecvInfo.size()));
205 1 : CHK_RET(HcclBatchSendRecvGroup(
206 : planner->sendRecvInfo.data(), planner->sendRecvInfo.size(), comm, planner->sendRecvMainStream));
207 : }
208 1 : HCCL_INFO("[doLaunches] take time [%lld]us.", DURATION_US(TIME_NOW() - startutime));
209 1 : if (planner->nTasksColl != 0) {
210 : /* 按数据量降序排列,V类算子各rank数据量可能不同,统一放到最后,保持原序 */
211 1 : auto sortedTasks = sortGroupTasks(planner->collTaskQueue);
212 1 : HCCL_INFO("Collectives sorted!");
213 1 : for (auto& taskColl : sortedTasks) {
214 0 : switch (taskColl.coll) {
215 0 : case HcclCMDType::HCCL_CMD_ALLGATHER:
216 0 : HCCL_INFO("AllGather, sendCount[%llu]", taskColl.sendCount);
217 0 : CHK_RET(HcclAllGatherInner(
218 : taskColl.sendbuff, taskColl.recvbuff, taskColl.sendCount, taskColl.sendType, taskColl.comm,
219 : taskColl.stream));
220 0 : break;
221 0 : case HcclCMDType::HCCL_CMD_REDUCE_SCATTER:
222 0 : HCCL_INFO("ReduceScatter, recvCount[%llu]", taskColl.recvCount);
223 0 : CHK_RET(HcclReduceScatterInner(
224 : taskColl.sendbuff, taskColl.recvbuff, taskColl.recvCount, taskColl.recvType, taskColl.op,
225 : taskColl.comm, taskColl.stream));
226 0 : break;
227 0 : case HcclCMDType::HCCL_CMD_ALLREDUCE:
228 0 : HCCL_INFO("AllReduce, sendCount[%llu]", taskColl.sendCount);
229 0 : CHK_RET(HcclAllReduceInner(
230 : taskColl.sendbuff, taskColl.recvbuff, taskColl.sendCount, taskColl.sendType, taskColl.op,
231 : taskColl.comm, taskColl.stream));
232 0 : break;
233 0 : case HcclCMDType::HCCL_CMD_BROADCAST:
234 0 : CHK_RET(HcclBroadcastInner(
235 : taskColl.sendbuff, taskColl.sendCount, taskColl.sendType, taskColl.root, taskColl.comm,
236 : taskColl.stream));
237 0 : break;
238 0 : case HcclCMDType::HCCL_CMD_ALLTOALL:
239 0 : CHK_RET(HcclAlltoAllInner(
240 : taskColl.sendbuff, taskColl.sendCount, taskColl.sendType, taskColl.recvbuff, taskColl.recvCount,
241 : taskColl.recvType, taskColl.comm, taskColl.stream));
242 0 : break;
243 0 : case HcclCMDType::HCCL_CMD_ALLTOALLV:
244 0 : CHK_RET(HcclAlltoAllVInner(
245 : taskColl.sendbuff, taskColl.sendCounts, taskColl.sdispls, taskColl.sendType, taskColl.recvbuff,
246 : taskColl.recvCounts, taskColl.rdispls, taskColl.recvType, taskColl.comm, taskColl.stream));
247 0 : break;
248 0 : case HcclCMDType::HCCL_CMD_ALLTOALLVC:
249 0 : CHK_RET(HcclAlltoAllVCInner(
250 : taskColl.sendbuff, taskColl.sendCounts, taskColl.sendType, taskColl.recvbuff, taskColl.recvType,
251 : taskColl.comm, taskColl.stream));
252 0 : break;
253 0 : case HcclCMDType::HCCL_CMD_REDUCE:
254 0 : CHK_RET(HcclReduceInner(
255 : taskColl.sendbuff, taskColl.recvbuff, taskColl.recvCount, taskColl.recvType, taskColl.op,
256 : taskColl.root, taskColl.comm, taskColl.stream));
257 0 : break;
258 0 : case HcclCMDType::HCCL_CMD_SCATTER:
259 0 : CHK_RET(HcclScatterInner(
260 : taskColl.sendbuff, taskColl.recvbuff, taskColl.recvCount, taskColl.recvType, taskColl.root,
261 : taskColl.comm, taskColl.stream));
262 0 : break;
263 0 : case HcclCMDType::HCCL_CMD_ALLGATHER_V:
264 0 : CHK_RET(HcclAllGatherVInner(
265 : taskColl.sendbuff, taskColl.sendCount, taskColl.recvbuff, taskColl.recvCounts, taskColl.rdispls,
266 : taskColl.recvType, taskColl.comm, taskColl.stream));
267 0 : break;
268 0 : case HcclCMDType::HCCL_CMD_REDUCE_SCATTER_V:
269 0 : CHK_RET(HcclReduceScatterVInner(
270 : taskColl.sendbuff, taskColl.sendCounts, taskColl.sdispls, taskColl.recvbuff, taskColl.recvCount,
271 : taskColl.sendType, taskColl.op, taskColl.comm, taskColl.stream));
272 0 : break;
273 0 : default:
274 0 : HCCL_ERROR("[doLaunches] not supported hcclFunc!");
275 0 : return HCCL_E_INTERNAL;
276 : }
277 : }
278 1 : }
279 1 : return HCCL_SUCCESS;
280 1 : }
281 :
282 1 : static HcclResult groupLaunch()
283 : { // 将各种通信域初始化/destroy的asyncJobs,在这里触发放到背景线程执行
284 1 : HCCL_INFO("[groupLaunch] entered");
285 :
286 1 : asyncJobLaunch();
287 1 : HCCL_DEBUG("[groupLaunch] asyncJobLaunch done");
288 2 : for (HcclComm comm : hcclGroupCommList) {
289 1 : doLaunches(comm);
290 : }
291 1 : HCCL_INFO("[groupLaunch] doLaunches done");
292 : // 流同步
293 2 : for (HcclComm comm : hcclGroupCommList) {
294 1 : hccl::hcclComm* hcclComm = static_cast<hccl::hcclComm*>(comm);
295 1 : std::shared_ptr<struct hcclKernelPlanner> planner = hcclComm->planner;
296 1 : for (auto it : planner->collStreams) {
297 0 : CHK_RET(hcclStreamSynchronize(it));
298 : }
299 1 : }
300 1 : HCCL_INFO("groupLauch Done!");
301 1 : return HCCL_SUCCESS;
302 : }
303 :
304 1 : inline void groupLocalResetJobState()
305 : {
306 : // hcclcomm中group相关的变量
307 2 : for (HcclComm comm : hcclGroupCommList) {
308 1 : hccl::hcclComm* hcclComm = static_cast<hccl::hcclComm*>(comm);
309 1 : hcclComm->planner = std::make_shared<hcclKernelPlanner>();
310 1 : hcclComm->SetGroupMode(false);
311 : }
312 1 : hcclGroupCommList.clear();
313 :
314 1 : return;
315 : }
316 :
317 1 : HcclResult HcclLegacyGroupEnd()
318 : {
319 1 : groupLaunch();
320 1 : HCCL_INFO("[GroupEnd] done groupLaunch");
321 1 : groupLocalResetJobState();
322 1 : HCCL_INFO("[GroupEnd] to the end");
323 1 : return HCCL_SUCCESS;
324 : }
325 :
326 0 : HcclResult HcclLegacyAsyncJobLaunch() { return asyncJobLaunch(); }
|