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 : #include "alg_template_multi_deter_pipeline.h"
11 : namespace hccl {
12 0 : MultiDeterPipeline::MultiDeterPipeline(const HcclDispatcher dispatcher) : AlgTemplateBase(dispatcher) {}
13 :
14 0 : MultiDeterPipeline::~MultiDeterPipeline() {}
15 :
16 0 : HcclResult MultiDeterPipeline::RunAsync() { return HCCL_SUCCESS; }
17 :
18 : // ReduceScatterDeterPipeline
19 0 : HcclResult MultiDeterPipeline::Prepare(
20 : HcomCollOpInfo* opInfo, DeviceMem& buffer, const u64 count, const u64 offset, const std::vector<Slice>& slices,
21 : const SubCommInfo& level0CommInfo, const SubCommInfo& level1CommInfo, Stream& mainStream,
22 : std::vector<Stream>& subStream, std::vector<std::shared_ptr<LocalNotify>>& notifyMain,
23 : std::vector<std::shared_ptr<LocalNotify>>& notifySub)
24 : {
25 0 : return HCCL_SUCCESS;
26 : }
27 :
28 : // AllReduceDeterPipeline
29 0 : HcclResult MultiDeterPipeline::Prepare(
30 : HcomCollOpInfo* opInfo, DeviceMem& inBuffer, DeviceMem& outBuffer, const u64 count,
31 : const std::vector<Slice>& slices, const SubCommInfo& level0CommInfo, const SubCommInfo& level1CommInfo,
32 : Stream& mainStream, std::vector<Stream>& subStream, std::vector<std::shared_ptr<LocalNotify>>& notifyMain,
33 : std::vector<std::shared_ptr<LocalNotify>>& notifySub)
34 : {
35 0 : return HCCL_SUCCESS;
36 : }
37 :
38 0 : HcclResult MultiDeterPipeline::MainWaitSub(u32 begin, u32 end)
39 : {
40 0 : for (u32 signalIndex = begin; signalIndex < end; signalIndex++) {
41 0 : CHK_RET(LocalNotify::Wait(mainStream_, dispatcher_, streamNotifyMain_[signalIndex], INVALID_VALUE_STAGE));
42 : }
43 0 : return HCCL_SUCCESS;
44 : }
45 :
46 0 : HcclResult MultiDeterPipeline::SubRecordMain(u32 begin, u32 end)
47 : {
48 0 : for (u32 streamIndex = begin; streamIndex < end; streamIndex++) {
49 0 : CHK_RET(LocalNotify::Post(subStreams_[streamIndex], dispatcher_, streamNotifyMain_[streamIndex], -1));
50 : }
51 0 : return HCCL_SUCCESS;
52 : }
53 :
54 0 : HcclResult MultiDeterPipeline::MainRecordSub(u32 begin, u32 end)
55 : {
56 0 : for (u32 signalIndex = begin; signalIndex < end; signalIndex++) {
57 0 : CHK_RET(LocalNotify::Post(mainStream_, dispatcher_, streamNotifySub_[signalIndex], -1));
58 : }
59 0 : return HCCL_SUCCESS;
60 : }
61 :
62 : // begin max = 7, end max = 11
63 0 : HcclResult MultiDeterPipeline::SubWaitMain(u32 begin, u32 end)
64 : {
65 0 : for (u32 streamIndex = begin; streamIndex < end; streamIndex++) {
66 0 : CHK_RET(LocalNotify::Wait(
67 : subStreams_[streamIndex], dispatcher_, streamNotifySub_[streamIndex], INVALID_VALUE_STAGE));
68 : }
69 0 : return HCCL_SUCCESS;
70 : }
71 :
72 0 : HcclResult MultiDeterPipeline::GetRemoteCclbufferDeviceMem(
73 : u32 inputSliceIndex, LINK link, u32 outputSliceIndex, DeviceMem& remoteMem)
74 : {
75 0 : return HCCL_SUCCESS;
76 : }
77 :
78 0 : HcclResult MultiDeterPipeline::GetLocalUserInDeviceMem(u32 rankIdInAllRanks, DeviceMem& locaMem)
79 : {
80 0 : return HCCL_SUCCESS;
81 : }
82 :
83 0 : HcclResult MultiDeterPipeline::GetLocalUserOutDeviceMem(u32 rankIdInAllRanks, DeviceMem& localMem)
84 : {
85 0 : return HCCL_SUCCESS;
86 : }
87 :
88 : HcclResult
89 0 : MultiDeterPipeline::GetLocalInCclbufferDeviceMem(u32 rankIdInAllRanks, DeviceMem& localMem, bool ifUseLastSize)
90 : {
91 0 : return HCCL_SUCCESS;
92 : }
93 :
94 : HcclResult
95 0 : MultiDeterPipeline::GetLocalOutCclbufferDeviceMem(u32 rankIdInAllRanks, DeviceMem& localMem, bool ifUseLastSize)
96 : {
97 0 : return HCCL_SUCCESS;
98 : }
99 :
100 0 : HcclResult MultiDeterPipeline::RunLocalCopy() { return HCCL_SUCCESS; }
101 :
102 0 : HcclResult MultiDeterPipeline::RunIntraAlltoallPreSync(u32 step) { return HCCL_SUCCESS; }
103 :
104 0 : HcclResult MultiDeterPipeline::RunIntraAlltoall(u32 step)
105 : {
106 0 : u32 recvServerId = GetPreServerIdByStep(step); // 从上一个收 2
107 0 : u32 sendServerId = GetNextServerIdByStep(step); // 发给发下一个 1
108 : // 机内alltoall full mesh收集数据,是为了收集机内第sendServerId整块的内存(包含intraRankId_块)
109 : // 该索引是为了计算第sendServerId整块内的内存序号块(0~intraRankId_-1)
110 0 : std::vector<u32> localUsrInIndex;
111 0 : for (u32 i = intraRankId_ + 1; i < intraRankSize_ + intraRankId_; ++i) {
112 0 : localUsrInIndex.push_back(i % intraRankSize_);
113 : }
114 0 : for (u32 i = 0; i < intraRankSize_ - 1; ++i) {
115 0 : u32 sendIntraRankId = GetNextIntraRankIdByStep(i + 1);
116 0 : LINK sendIntraLink = intraLinks_[sendIntraRankId];
117 0 : DeviceMem srcMem;
118 0 : DeviceMem dstMem;
119 : // 从usrin收集发给下一个cclbufer的数据, 收集的所有数据需要发送给机间序号为sendServerId的server
120 0 : u32 needSendInputndex = GetRankIdx(sendServerId, localUsrInIndex[i]);
121 : // 发送数据到cclbufer 索引为[intraRankId_, localUsrInIndex[i]]
122 0 : u32 recvIntraRankIdx = alltoallRecvBlockIdxMap_[intraRankId_][localUsrInIndex[i]];
123 0 : u32 recvCclbufferIndex = GetRankIdx(recvServerId, recvIntraRankIdx);
124 0 : CHK_RET(GetLocalUserInDeviceMem(needSendInputndex, srcMem));
125 0 : CHK_RET(GetRemoteCclbufferDeviceMem(needSendInputndex, sendIntraLink, recvCclbufferIndex, dstMem));
126 : // 发送给机内rank索引 [serverId_, sendIntraRankId]
127 0 : u32 remoteUserRank = GetRankIdx(serverId_, sendIntraRankId);
128 : // SDMA copy write语义,因为只能写到cclbuffer
129 0 : CHK_RET(HcclD2DMemcpyAsync(
130 : dispatcher_, dstMem, srcMem, subStreams_[i], remoteUserRank, sendIntraLink->GetLinkType()));
131 0 : CHK_RET(sendIntraLink->TxDataSignal(subStreams_[i]));
132 0 : CHK_RET(sendIntraLink->RxDataSignal(subStreams_[i]));
133 0 : HCCL_DEBUG("[%s] intra-server SDMA send, intraRank: [%u] -> [%u]", __func__, intraRankId_, sendIntraRankId);
134 0 : HCCL_DEBUG(
135 : "[%s] intra-server SDMA send, mem: inputMem[%u, %u] -> cclbuffer[%u, %u]; cclbufferNo[%u] -> [%u]",
136 : __func__, sendServerId, localUsrInIndex[i], recvServerId, recvIntraRankIdx, needSendInputndex,
137 : recvCclbufferIndex);
138 0 : }
139 0 : HCCL_INFO("[%s] intra-server step[%u] run alltoall success", __func__, step);
140 0 : return HCCL_SUCCESS;
141 0 : }
142 :
143 0 : HcclResult MultiDeterPipeline::GroupTasksByStream(
144 : u32 activeCount, const std::vector<bool>& isReduceBlock, u32 retIndex,
145 : std::vector<std::vector<std::vector<std::pair<u32, u32>>>>& batchStreamTasks, // 输出:批次→流→任务
146 : std::vector<bool>& processed, std::vector<u32>& origIdxMap, u32& newActiveCount)
147 : {
148 0 : batchStreamTasks.clear(); // 清空批次任务
149 0 : processed.assign(activeCount, false);
150 0 : newActiveCount = 0;
151 :
152 0 : const u32 mergeStep = 2;
153 0 : const u32 batchSize = MAX_REDUCE_STREAM_NUM;
154 0 : const u32 totalGroups = (activeCount + mergeStep - 1) / mergeStep;
155 0 : const u32 batchNum = (totalGroups + batchSize - 1) / batchSize;
156 :
157 : // 逐批生成流任务
158 0 : for (u32 batch = 0; batch < batchNum; batch++) {
159 : // 初始化当前批次的流任务(MAX_REDUCE_STREAM_NUM条流)
160 0 : std::vector<std::vector<std::pair<u32, u32>>> streamTasks(MAX_REDUCE_STREAM_NUM);
161 0 : u32 startGroup = batch * batchSize;
162 0 : u32 endGroup = std::min((batch + 1) * batchSize, totalGroups);
163 :
164 : // 处理当前批次的分组
165 0 : for (u32 group = startGroup; group < endGroup; group++) {
166 0 : u32 idx0 = group * mergeStep;
167 0 : u32 idx1 = idx0 + 1;
168 0 : if (idx1 >= activeCount) {
169 0 : processed[idx0] = false;
170 0 : continue;
171 : }
172 :
173 : // 选择dst/src(原有优先级逻辑不变)
174 : u32 dstIdx, srcIdx;
175 0 : if (origIdxMap[idx0] == retIndex) { // 当前idx0的原始索引是目标块,强制为dst
176 0 : dstIdx = idx0;
177 0 : srcIdx = idx1;
178 0 : } else if (origIdxMap[idx1] == retIndex) { // 当前idx1的原始索引是目标块,强制为dst
179 0 : dstIdx = idx1;
180 0 : srcIdx = idx0;
181 0 : } else if (isReduceBlock[idx0] && !isReduceBlock[idx1]) {
182 0 : dstIdx = idx0;
183 0 : srcIdx = idx1;
184 0 : } else if (!isReduceBlock[idx0] && isReduceBlock[idx1]) {
185 0 : dstIdx = idx1;
186 0 : srcIdx = idx0;
187 : } else {
188 0 : dstIdx = std::max(idx0, idx1);
189 0 : srcIdx = std::min(idx0, idx1);
190 : }
191 :
192 : // 分配到当前批次的流任务中
193 0 : u32 batchInnerGroupIdx = group - startGroup;
194 0 : u32 streamId = batchInnerGroupIdx % MAX_REDUCE_STREAM_NUM;
195 0 : streamTasks[streamId].emplace_back(srcIdx, dstIdx);
196 :
197 0 : processed[srcIdx] = true;
198 0 : processed[dstIdx] = false;
199 :
200 0 : HCCL_DEBUG(
201 : "[%s] batch[%u] group[%u] merge src[%u] -> dst[%u] on stream[%u]", __func__, batch, group, srcIdx,
202 : dstIdx, streamId + reduceStreamBegin_);
203 : }
204 :
205 : // 将当前批次的流任务加入总批次列表
206 0 : batchStreamTasks.push_back(streamTasks);
207 0 : }
208 :
209 0 : newActiveCount = std::count(processed.begin(), processed.end(), false);
210 0 : return HCCL_SUCCESS;
211 : }
212 :
213 0 : HcclResult MultiDeterPipeline::BatchPostNotifyForStreams(
214 : const std::vector<std::vector<std::pair<u32, u32>>>& streamTasks, bool isStartPhase, bool useMainStream)
215 : {
216 0 : return HCCL_SUCCESS;
217 : }
218 :
219 0 : HcclResult MultiDeterPipeline::ExecuteStreamTasks(
220 : const std::vector<std::vector<std::pair<u32, u32>>>& streamTasks, const std::vector<DeviceMem>& validMem,
221 : std::vector<u32>& origIdxMap, bool useMainStream)
222 : {
223 0 : for (u32 s = 0; s < MAX_REDUCE_STREAM_NUM; s++) {
224 0 : if (streamTasks[s].empty())
225 0 : continue;
226 :
227 0 : u32 streamIdx = reduceStreamBegin_ + s;
228 0 : Stream& subStream = subStreams_[streamIdx];
229 0 : Stream& stream = useMainStream ? mainStream_ : subStream;
230 :
231 0 : for (const auto& task : streamTasks[s]) {
232 0 : u32 srcIdx = task.first;
233 0 : u32 dstIdx = task.second;
234 0 : const DeviceMem& dstMem = validMem[dstIdx];
235 0 : const DeviceMem& srcMem = validMem[srcIdx];
236 0 : u64 count = srcMem.size() / unitSize_;
237 0 : CHK_RET(HcclReduceAsync(
238 : dispatcher_, srcMem.ptr(), count, dataType_, reductionOp_, stream, dstMem.ptr(), INVALID_VALUE_RANKID,
239 : LinkType::LINK_ONCHIP, INLINE_REDUCE_BIT));
240 0 : HCCL_DEBUG(
241 : "[%s] stream[%u] execute task: merge src[%u] -> dst[%u], origSrc[%u] -> origDst[%u]", __func__,
242 : useMainStream ? 0 : streamIdx, srcIdx, dstIdx, origIdxMap[srcIdx], origIdxMap[dstIdx]);
243 : }
244 : }
245 0 : return HCCL_SUCCESS;
246 : }
247 :
248 0 : void MultiDeterPipeline::CompressActiveSet(
249 : std::vector<DeviceMem>& validMem, std::vector<bool>& isReduceBlock, std::vector<u32>& origIdxMap,
250 : const std::vector<bool>& processed, u32& trackedTargetIdx, const u32 origRetIndex)
251 : {
252 0 : std::vector<DeviceMem> newValidMem;
253 0 : std::vector<bool> newIsReduceBlock;
254 0 : std::vector<u32> newOrigIdxMap;
255 0 : u32 newTrackedTargetIdx = 0;
256 0 : bool foundTarget = false;
257 :
258 : // 保留未处理的块(processed=false,即dst块)
259 0 : for (u32 i = 0; i < validMem.size(); i++) {
260 0 : if (!processed[i]) { // 仅保留dst块,移除src块(processed=true)
261 0 : newValidMem.push_back(validMem[i]);
262 0 : newIsReduceBlock.push_back(isReduceBlock[i]);
263 0 : newOrigIdxMap.push_back(origIdxMap[i]);
264 :
265 : // 追踪原始目标块(retIndex)的新索引
266 0 : if (!foundTarget && origIdxMap[i] == origRetIndex) {
267 0 : newTrackedTargetIdx = newValidMem.size() - 1;
268 0 : foundTarget = true;
269 : }
270 : }
271 : }
272 :
273 0 : validMem.swap(newValidMem);
274 0 : isReduceBlock.swap(newIsReduceBlock);
275 0 : origIdxMap.swap(newOrigIdxMap);
276 : // 未找到目标块时,默认指向最后一个块
277 0 : trackedTargetIdx = foundTarget ? newTrackedTargetIdx : (validMem.empty() ? 0 : validMem.size() - 1);
278 0 : HCCL_DEBUG(
279 : "[%s] compressed: old size[%llu], new size[%llu], trackedTargetIdx[%u]", __func__, processed.size(),
280 : validMem.size(), trackedTargetIdx);
281 0 : }
282 :
283 0 : HcclResult MultiDeterPipeline::LocalReduce(
284 : std::vector<DeviceMem>& reduceMem, std::vector<bool>& isReduceBlock, u32 retIndex, bool useMainStream)
285 : {
286 0 : const u32 totalBlockCount = reduceMem.size();
287 : // 校验1:容器大小匹配 + retIndex越界
288 0 : if (reduceMem.size() != isReduceBlock.size() || retIndex >= totalBlockCount) {
289 0 : HCCL_ERROR(
290 : "[%s] Invalid param (size mismatch: %llu vs %llu, retIndex: %u >= %u)", __func__, reduceMem.size(),
291 : isReduceBlock.size(), retIndex, totalBlockCount);
292 0 : return HCCL_E_PARA;
293 : }
294 : // 校验2:目标内存块有效
295 0 : const DeviceMem& targetCCLBuffer = reduceMem[retIndex];
296 0 : if (targetCCLBuffer.ptr() == nullptr || targetCCLBuffer.size() == 0) {
297 0 : HCCL_ERROR(
298 : "[%s] Target CCLBuffer invalid (ptr: %p, size: %llu)", __func__, targetCCLBuffer.ptr(),
299 : targetCCLBuffer.size());
300 0 : return HCCL_E_MEMORY;
301 : }
302 :
303 0 : std::vector<DeviceMem> validMem = std::move(reduceMem); // 外层不使用reduceMem
304 0 : std::vector<bool> validIsReduceBlock = std::move(isReduceBlock);
305 : // 1. 动态追踪目标块索引 2. 原索引→当前索引的映射表
306 0 : std::vector<u32> origIdxMap(validMem.size());
307 0 : for (size_t i = 0; i < origIdxMap.size(); ++i) {
308 0 : origIdxMap[i] = i;
309 : }
310 :
311 0 : if (validMem.size() == 1) {
312 0 : HCCL_ERROR("[%s] validMem size is one, only target block valid", __func__);
313 0 : return HCCL_E_PARA;
314 : }
315 :
316 0 : u32 trackedTargetIdx = retIndex;
317 0 : u32 activeCount = validMem.size();
318 0 : const u32 origRetIndex = retIndex;
319 :
320 0 : while (activeCount > 1) {
321 0 : std::vector<std::vector<std::vector<std::pair<u32, u32>>>> batchStreamTasks; // 批次→流→任务
322 0 : std::vector<bool> processed(activeCount, false);
323 0 : u32 newActiveCount = 0;
324 :
325 : // 2.1 生成按批次组织的流任务
326 0 : CHK_RET(GroupTasksByStream(
327 : activeCount, validIsReduceBlock, origRetIndex, batchStreamTasks, processed, origIdxMap, newActiveCount));
328 :
329 : // 2.2 逐批执行流任务(串行处理每批)
330 0 : for (const auto& streamTasks : batchStreamTasks) {
331 : // a. 执行start phase notify(仅处理有任务的流)
332 0 : CHK_RET(BatchPostNotifyForStreams(streamTasks, true, useMainStream));
333 : // b. 执行当前批次的流任务(src→dst归约)
334 0 : CHK_RET(ExecuteStreamTasks(streamTasks, validMem, origIdxMap, useMainStream));
335 : // c. 执行sync phase notify(等待当前批次完成)
336 0 : CHK_RET(BatchPostNotifyForStreams(streamTasks, false, useMainStream));
337 : }
338 : // 2.3 压缩活跃块(移除src块,保留dst块)
339 0 : CompressActiveSet(validMem, validIsReduceBlock, origIdxMap, processed, trackedTargetIdx, origRetIndex);
340 0 : activeCount = validMem.size();
341 0 : HCCL_DEBUG("[LocalReduce] round done: activeCount=%u -> %u", newActiveCount, activeCount);
342 0 : }
343 0 : HCCL_DEBUG(
344 : "[%s] Local reduce success (merge to retIndex[%u] CCLBuffer, final tracked idx[%u])", __func__, retIndex,
345 : trackedTargetIdx);
346 0 : return HCCL_SUCCESS;
347 0 : }
348 :
349 0 : HcclResult MultiDeterPipeline::RunIntraLocalReduce(u32 step) { return HCCL_SUCCESS; }
350 :
351 0 : HcclResult MultiDeterPipeline::RunInterSend(u32 step) { return HCCL_SUCCESS; }
352 :
353 0 : HcclResult MultiDeterPipeline::RunFinalReduce() { return HCCL_SUCCESS; }
354 :
355 0 : HcclResult MultiDeterPipeline::AlltoallSync(u32 step, bool isStartPhase) { return HCCL_SUCCESS; }
356 :
357 0 : HcclResult MultiDeterPipeline::LocalReduceSync(u32 step, bool isStartPhase) { return HCCL_SUCCESS; }
358 :
359 0 : HcclResult MultiDeterPipeline::AlltoallLocalReduceSync(u32 step, bool isStartPhase)
360 : {
361 0 : bool alltoallStep = (step < allSteps_);
362 0 : bool localReduceStep = (step > 1 && step < allSteps_ + 1);
363 0 : if (alltoallStep) {
364 0 : CHK_RET(AlltoallSync(step, isStartPhase));
365 : }
366 0 : if (localReduceStep) {
367 0 : CHK_RET(LocalReduceSync(step, isStartPhase));
368 : }
369 0 : return HCCL_SUCCESS;
370 : }
371 :
372 0 : HcclResult MultiDeterPipeline::RunAsyncLocalReduceSerial()
373 : {
374 0 : HCCL_INFO(
375 : "[MultiDeterPipeline] run begin: rank[%u] ranksize[%u] inputMem[%p] outputMem[%p]", userRank_, userRankSize_,
376 : usrInMemPtr_, usrOutMemPtr_);
377 0 : CHK_SMART_PTR_NULL(dispatcher_);
378 : // 以机内8卡为例主流 + 从流 = 1 + 7 + 4 = 12
379 0 : allSteps_ = serverSize_;
380 : // #1 机间发送,#2 local reduce,#3 机内alltoall
381 : // 以机间的pairwise来分步,#n表示pairwise的第n步
382 0 : for (u32 step = 1; step < allSteps_ + 1; step++) {
383 0 : HCCL_DEBUG(
384 : "[%s] userRank[%u], intraRankId[%u], serverId[%u], step[%u/%u] begin", __func__, userRank_, intraRankId_,
385 : serverId_, step, allSteps_);
386 : // alltoall 主从流同步+前同步+拉齐
387 0 : CHK_RET(AlltoallSync(step, true));
388 0 : if (step < allSteps_ - 1) {
389 0 : CHK_RET(RunIntraAlltoallPreSync(step));
390 : }
391 0 : if (step == allSteps_ - 1) {
392 0 : CHK_RET(RunIntraAlltoallPreSync(0));
393 : }
394 0 : if (step == 1) {
395 0 : CHK_RET(RunLocalCopy());
396 : }
397 0 : if (step > 1) {
398 0 : CHK_RET(RunInterSend(step - 1));
399 : }
400 : // #1 机内RS alltoall + local reduce串行
401 0 : if (step < allSteps_) {
402 0 : CHK_RET(RunIntraAlltoall(step));
403 0 : CHK_RET(AlltoallSync(step, false));
404 0 : CHK_RET(LocalReduceSync(step, true));
405 0 : CHK_RET(RunIntraLocalReduce(step));
406 0 : CHK_RET(LocalReduceSync(step, false));
407 : } else {
408 : // #0 机内RS alltoall + local reduce串行,#0不需要向其他机发送数据
409 0 : CHK_RET(RunIntraAlltoall(0));
410 0 : CHK_RET(AlltoallSync(0, false));
411 0 : CHK_RET(LocalReduceSync(0, true));
412 0 : CHK_RET(RunIntraLocalReduce(0));
413 0 : CHK_RET(LocalReduceSync(0, false));
414 : }
415 : }
416 : // 总local reduce
417 0 : CHK_RET(RunFinalReduce());
418 0 : HCCL_INFO("[MultiDeterPipeline] MultiDeterPipeline success userRank[%u] ", userRank_);
419 0 : return HCCL_SUCCESS;
420 : }
421 :
422 : // 每个server内首先要进行alltoall full mesh收集数据,再进行机内local reduce,最后发送给指定server
423 0 : HcclResult MultiDeterPipeline::RunAsyncReduceScatterPipeline()
424 : {
425 0 : constexpr u64 HCCL_MEDIUM_COUNT_2_MB = 2 * 1024 * 1024;
426 : // 2机或者数据量小于2MB走localreduce串行算法
427 0 : if (serverSize_ <= LOCAL_REDUCE_SERIIAL_ALG_SERVER_NUM || GetLocalReduceSerialThresh() < HCCL_MEDIUM_COUNT_2_MB) {
428 0 : CHK_RET(RunAsyncLocalReduceSerial());
429 0 : return HCCL_SUCCESS;
430 : }
431 0 : HCCL_INFO("[MultiDeterPipeline] [%s] begin, userRank[%u]", __func__, userRank_);
432 : // 以机内8卡为例主流 + 从流 = 1 + 7 + 4 = 12
433 : // pairwise总共需要serverSize_步,#1机内alltoall只能自己执行,无法和其他步骤并行,所以总共需要serverSize_ + 1个步骤
434 0 : allSteps_ = serverSize_ + 1;
435 : // #1 机间发送,#2 local reduce,#3 机内alltoall
436 : // 以机间的pairwise来分步,#n表示pairwise的第n步
437 0 : for (u32 step = 1; step < allSteps_ + 1; step++) {
438 0 : HCCL_DEBUG(
439 : "[%s] userRank[%u], intraRankId[%u], serverId[%u], step[%u/%u] begin", __func__, userRank_, intraRankId_,
440 : serverId_, step, allSteps_);
441 0 : CHK_RET(AlltoallLocalReduceSync(step, true));
442 0 : if (step == 1) {
443 0 : CHK_RET(RunLocalCopy());
444 : }
445 0 : if (step < allSteps_ - 1) {
446 0 : CHK_RET(RunIntraAlltoallPreSync(step));
447 : }
448 0 : if (step == allSteps_ - 1) {
449 0 : CHK_RET(RunIntraAlltoallPreSync(0));
450 : }
451 0 : if (step > STEP_OFFSET_TWO) {
452 0 : CHK_RET(RunInterSend(step - STEP_OFFSET_TWO));
453 : }
454 : // allSteps_最小为3
455 0 : if (step > 1 && step < allSteps_) {
456 0 : CHK_RET(RunIntraLocalReduce(step - 1));
457 : }
458 : // 最后一步,需要额外执行 #0 机内RS local reduce,#0不需要向其他机发送数据
459 0 : if (step == allSteps_) {
460 0 : CHK_RET(RunIntraLocalReduce(0));
461 : }
462 0 : if (step < allSteps_ - 1) {
463 0 : CHK_RET(RunIntraAlltoall(step));
464 : }
465 0 : if (step == allSteps_ - 1) {
466 : // #0 机内RS alltoall
467 0 : CHK_RET(RunIntraAlltoall(0));
468 : }
469 0 : CHK_RET(AlltoallLocalReduceSync(step, false));
470 : }
471 : // 总local reduce
472 0 : CHK_RET(RunFinalReduce());
473 0 : HCCL_INFO("[MultiDeterPipeline] [%s] end, userRank[%u]", __func__, userRank_);
474 0 : return HCCL_SUCCESS;
475 : }
476 :
477 : // 遍历所有发送方rank和发送方块索引,预计算映射关系,目的是构造出下属矩阵
478 : // srcRank\srcBlockIdx 0 1 2
479 : // 0 MAX 0 0
480 : // 1 0 MAX 1
481 : // 2 1 1 MAX
482 : // 例1,发送方是rank0,接收端是rank1,那么接收端非自身 rank 列表为[0, 2],那么rank0的索引为0,所以就发到rank1的第0块内存
483 : // 例2,发送方是rank0,接收端是rank2,那么接收端非自身 rank 列表为[0, 1],那么rank0的索引为0,所以就发到rank2的第0块内存
484 : // 例3,发送方是rank1,接收端是rank2,那么接收端非自身 rank 列表为[0, 1],那么rank1的索引为1,所以就发到rank2的第1块内存
485 0 : void MultiDeterPipeline::InitAlltoallRecvBlockIdxMap()
486 : {
487 0 : const u32 rankSize = intraRankSize_;
488 0 : alltoallRecvBlockIdxMap_.resize(rankSize, std::vector<u32>(rankSize, UINT32_MAX));
489 0 : if (rankSize <= 1) {
490 0 : return;
491 : }
492 : // 规则一:接收端的块索引 = 发送方 rank 在「接收端非自身 rank 列表」中的索引
493 0 : std::vector<std::vector<u32>> dstRankToRankIndex(rankSize, std::vector<u32>(rankSize, UINT32_MAX));
494 0 : for (u32 dstRank = 0; dstRank < rankSize; ++dstRank) {
495 0 : for (u32 srcRank = 0; srcRank < rankSize; ++srcRank) {
496 0 : if (srcRank == dstRank) {
497 0 : continue;
498 : }
499 : // 若 srcRank < dstRank:索引 = srcRank, 若 srcRank > dstRank:索引 = srcRank - 1
500 0 : const u32 idx = (srcRank < dstRank) ? srcRank : (srcRank - 1);
501 0 : dstRankToRankIndex[dstRank][srcRank] = idx;
502 : }
503 : }
504 0 : for (u32 srcRank = 0; srcRank < rankSize; ++srcRank) {
505 0 : for (u32 srcBlockIdx = 0; srcBlockIdx < rankSize; ++srcBlockIdx) {
506 : // 规则二:发送方块索引 = 接收端 rank
507 0 : const u32 dstRank = srcBlockIdx;
508 : // 跳过自身发送的无效场景
509 0 : if (srcRank == dstRank) {
510 0 : continue;
511 : }
512 : // 查找发送方rank在列表中的索引,存入映射表
513 0 : const u32 dstBlockIdx = dstRankToRankIndex[dstRank][srcRank];
514 0 : alltoallRecvBlockIdxMap_[srcRank][srcBlockIdx] = dstBlockIdx;
515 0 : HCCL_DEBUG(
516 : "[%s] srcRank[%u], srcBlockIdx[%u] -> dstRank[%u], dstBlockIdx[%u]", __func__, srcRank, srcBlockIdx,
517 : dstRank, dstBlockIdx);
518 : }
519 : }
520 0 : }
521 :
522 0 : HcclResult MultiDeterPipeline::PrepareTopoInfo(const SubCommInfo& level0CommInfo, const SubCommInfo& level1CommInfo)
523 : {
524 0 : serverSize_ = level1CommInfo.localRankSize;
525 0 : CHK_PRT_RET(
526 : serverSize_ < MIN_SERVER_NUM,
527 : HCCL_ERROR("[%s] Unexpected inter rank size[%u], which should >= 2.", __func__, serverSize_), HCCL_E_PARA);
528 :
529 0 : intraRankSize_ = level0CommInfo.localRankSize;
530 0 : CHK_PRT_RET(
531 : intraRankSize_ < MIN_INTRA_RANK_NUM,
532 : HCCL_ERROR("[%s] Unexpected intra rank size[%u], which should >= 3.", __func__, intraRankSize_), HCCL_E_PARA);
533 0 : intraRankId_ = level0CommInfo.localRank;
534 0 : serverId_ = level1CommInfo.localRank;
535 0 : userRankSize_ = intraRankSize_ * serverSize_;
536 0 : userRank_ = intraRankId_ + serverId_ * intraRankSize_;
537 :
538 0 : intraLinks_ = level0CommInfo.links; // 节点内
539 0 : serverLinks_ = level1CommInfo.links; // 节点间
540 0 : HCCL_INFO(
541 : "[%s] opInfo: dataType[%u], unitSize[%u], memSliceSize[%u], usrInMem[%p], usrOutMem[%p], reductionOp[%u]",
542 : __func__, dataType_, unitSize_, memSliceSize_, usrInMemPtr_, usrOutMemPtr_, reductionOp_);
543 0 : HCCL_INFO(
544 : "[%s] topoInfo: userRank[%u], intraRankId[%u], intraRankSize[%u], serverId[%u], interRankSize[%u]", __func__,
545 : userRank_, intraRankId_, intraRankSize_, serverId_, serverSize_);
546 0 : HCCL_INFO(
547 : "[%s] topoInfo: severLinksNum[%zu], intraLinksNum[%zu]", __func__, serverLinks_.size(), intraLinks_.size());
548 0 : InitAlltoallRecvBlockIdxMap();
549 0 : return HCCL_SUCCESS;
550 : }
551 : } // namespace hccl
|