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 "all_gather_mesh_atomic.h"
12 : #include "alg_template_register.h"
13 :
14 : namespace hccl {
15 0 : AllGatherMeshAtomic::AllGatherMeshAtomic(const HcclDispatcher dispatcher)
16 0 : : AllGatherMesh(dispatcher)
17 0 : {}
18 :
19 0 : AllGatherMeshAtomic::~AllGatherMeshAtomic() {}
20 :
21 0 : HcclResult AllGatherMeshAtomic::RunAllGather(const std::vector<LINK> &links, const std::vector<Slice> &outputSlices,
22 : const std::vector<Slice> &inputSlices)
23 : {
24 0 : for (u32 round = 1; round < interRankSize_; round++) {
25 0 : u32 dstRank = BackwardRank(interRank_, interRankSize_, round);
26 0 : Stream& subStream = (round == interRankSize_ - 1) ? stream_ : meshStreams_[round - 1];
27 0 : CHK_RET(links[dstRank]->TxAck(subStream));
28 0 : CHK_RET(links[dstRank]->RxAck(subStream));
29 : }
30 :
31 0 : for (u32 round = 1; round < interRankSize_; round++) {
32 0 : u32 dstRank = BackwardRank(interRank_, interRankSize_, round);
33 0 : Stream& subStream = (round == interRankSize_ - 1) ? stream_ : meshStreams_[round - 1];
34 0 : profilerInput_.streamID = subStream.id();
35 0 : profilerInput_.planeID = round - 1;
36 0 : profilerInput_.step = HCCL_EXEC_STEP_NOT_SET;
37 :
38 0 : if (round == interRankSize_ - 1) {
39 0 : for (u32 signalIndex = 0; signalIndex < interRankSize_ - 2; signalIndex++) { // rankSize-2: stream num
40 0 : CHK_RET(LocalNotify::Wait(subStream, dispatcher_, (*meshSignal_)[signalIndex], profilerInput_.stage));
41 : }
42 : // 为子图增加一个从stream到主stream的附着点
43 0 : DeviceMem src = DeviceMem::create(inputMem_.ptr(), 0);
44 0 : DeviceMem dst = DeviceMem::create(outputMem_.ptr(), 0);
45 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, stream_));
46 0 : for (u32 signalIndex = 0; signalIndex < interRankSize_ - 2; signalIndex++) { // rankSize-2: stream num
47 0 : CHK_RET(LocalNotify::Post(subStream, dispatcher_, (*meshSignalAux_)[signalIndex],
48 : profilerInput_.stage));
49 : }
50 0 : } else {
51 0 : u32 signalIndex = round - 1;
52 0 : CHK_RET(LocalNotify::Post(subStream, dispatcher_, (*meshSignal_)[signalIndex],
53 : profilerInput_.stage));
54 0 : CHK_RET(LocalNotify::Wait(subStream, dispatcher_, (*meshSignalAux_)[signalIndex], profilerInput_.stage));
55 : }
56 : // 本rank要收数据
57 0 : void *srcMemPtr = nullptr;
58 : // 从对端的input内存拿数据,input==output也没有关系
59 0 : CHK_RET(links[dstRank]->GetRemoteMem(UserMemType::OUTPUT_MEM, &srcMemPtr));
60 0 : DeviceMem srcDevMem(static_cast<s8 *>(srcMemPtr) + baseOffset_ + inputSlices[dstRank].offset,
61 0 : inputSlices[dstRank].size);
62 0 : DeviceMem dstDevMem = outputMem_.range(outputSlices[dstRank].offset, outputSlices[dstRank].size);
63 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dstDevMem, srcDevMem, subStream,
64 : links[dstRank]->GetRemoteRank(), links[dstRank]->GetLinkType()));
65 0 : CHK_RET(links[dstRank]->TxDataSignal(subStream));
66 0 : CHK_RET(links[dstRank]->RxDataSignal(subStream));
67 0 : }
68 :
69 0 : CHK_RET(AlgTemplateBase::ExecEmptyTask(inputMem_, outputMem_, stream_, dispatcher_));
70 0 : return HCCL_SUCCESS;
71 : }
72 : REGISTER_TEMPLATE(TemplateType::TEMPLATE_ALL_GATHER_MESH_ATOMIC, AllGatherMeshAtomic);
73 : } // namespace hccl
|