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_hccs_sio.h"
12 : #include "alg_template_register.h"
13 :
14 : namespace hccl {
15 0 : AllGatherHccsSio::AllGatherHccsSio(const HcclDispatcher dispatcher) : AlgTemplateBase(dispatcher) {}
16 :
17 0 : AllGatherHccsSio::~AllGatherHccsSio() {}
18 :
19 0 : HcclResult AllGatherHccsSio::Prepare(
20 : SubCommInfo& outerCommInfoHccs, SubCommInfo& outerCommInfoSio, DeviceMem& usrInMem, DeviceMem& usrOutMem, u64 count,
21 : const HcclDataType dataType, const Stream& mainStream, std::vector<Stream>& meshStreams,
22 : std::vector<std::shared_ptr<LocalNotify>>& meshSignal, std::vector<std::shared_ptr<LocalNotify>>& meshSignalAux,
23 : u32 userRank, HcomCollOpInfo* opInfo)
24 : {
25 0 : inputMem_ = usrInMem;
26 0 : outputMem_ = usrOutMem;
27 0 : stream_ = mainStream;
28 0 : meshStreams_ = meshStreams;
29 0 : meshSignal_ = meshSignal;
30 0 : meshSignalAux_ = meshSignalAux;
31 0 : userRank_ = userRank;
32 0 : dataType_ = dataType;
33 0 : dataBytes_ = count * SIZE_TABLE[dataType];
34 0 : count_ = count;
35 0 : outerCommInfoHccs_ = outerCommInfoHccs;
36 0 : outerCommInfoSio_ = outerCommInfoSio;
37 0 : opInfo_ = opInfo;
38 0 : totalDataBytes_ = opInfo->count * SIZE_TABLE[dataType_];
39 0 : return HCCL_SUCCESS;
40 : }
41 :
42 : // 主流所有从流
43 0 : HcclResult AllGatherHccsSio::NotifySubStreamStart()
44 : {
45 0 : for (u32 streamIndex = 0; streamIndex < meshStreams_.size(); streamIndex++) {
46 0 : CHK_RET(LocalNotify::Post(stream_, dispatcher_, meshSignalAux_[streamIndex], INVALID_VALUE_STAGE));
47 0 : CHK_RET(LocalNotify::Wait(
48 : meshStreams_[streamIndex], dispatcher_, meshSignalAux_[streamIndex], INVALID_VALUE_STAGE));
49 : }
50 0 : return HCCL_SUCCESS;
51 : }
52 :
53 0 : HcclResult AllGatherHccsSio::WaitSubStreamFinish()
54 : {
55 0 : for (u32 streamIndex = 0; streamIndex < meshStreams_.size(); streamIndex++) {
56 0 : CHK_RET(
57 : LocalNotify::Post(meshStreams_[streamIndex], dispatcher_, meshSignal_[streamIndex], INVALID_VALUE_STAGE));
58 0 : CHK_RET(LocalNotify::Wait(stream_, dispatcher_, meshSignal_[streamIndex], INVALID_VALUE_STAGE));
59 : }
60 0 : return HCCL_SUCCESS;
61 : }
62 :
63 : HcclResult
64 0 : AllGatherHccsSio::RunInterDie(const u32 dieRankId, const std::vector<LINK>& links, const u32 srcDMAMemSliceId)
65 : {
66 : // 检查链接是否为空
67 0 : CHK_SMART_PTR_NULL(links[dieRankId]);
68 :
69 : // 获取远程内存指针
70 0 : void* remDMAMemPtr = nullptr;
71 0 : CHK_RET(links[dieRankId]->GetRemoteMem(UserMemType::INPUT_MEM, &remDMAMemPtr));
72 :
73 : // 确定需要传输的数据部分(上半部分或下半部分)
74 0 : u64 dataPartOffset = dieRankId * dataBytes_;
75 0 : u64 dataPartSize = count_ / 2 * SIZE_TABLE[dataType_];
76 :
77 0 : DeviceMem locDieDst;
78 0 : DeviceMem srcDieMem;
79 :
80 : // 定义本地目标内存和远程源内存
81 0 : if (srcDMAMemSliceId == 0) {
82 0 : locDieDst = dmaMem_[1].range(dataPartOffset, dataPartSize);
83 0 : srcDieMem = DeviceMem::create(static_cast<u8*>(remDMAMemPtr), dataPartSize);
84 : } else {
85 0 : locDieDst = dmaMem_[1].range(dataPartOffset + dataPartSize, dataBytes_ - dataPartSize);
86 0 : srcDieMem = DeviceMem::create(static_cast<u8*>(remDMAMemPtr) + dataPartSize, dataBytes_ - dataPartSize);
87 : }
88 :
89 0 : HCCL_INFO(
90 : "RunInterDie: dieRankId[%d], locDieDst ptr[%p], locDieDst size[%ld], remDMAMemPtr[%p]", dieRankId,
91 : locDieDst.ptr(), locDieDst.size(), remDMAMemPtr);
92 : // 执行异步内存复制
93 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, locDieDst, srcDieMem, meshStreams_[srcDMAMemSliceId]));
94 :
95 0 : return HCCL_SUCCESS;
96 0 : }
97 :
98 : HcclResult
99 0 : AllGatherHccsSio::RunInterDieOpBase(const u32 dieRankId, const std::vector<LINK>& links, const u32 srcDMAMemSliceId)
100 : {
101 : // 检查链接是否为空
102 0 : CHK_SMART_PTR_NULL(links[dieRankId]);
103 :
104 : // 获取远程CCLin内存指针
105 0 : void* remCCLMemPtr = nullptr;
106 0 : CHK_RET(links[dieRankId]->GetRemoteMem(UserMemType::INPUT_MEM, &remCCLMemPtr));
107 :
108 : // 确定需要传输的数据部分(上半部分或下半部分)
109 0 : u64 dataPartSize = count_ / 2 * SIZE_TABLE[dataType_];
110 :
111 0 : DeviceMem locDieDst;
112 0 : DeviceMem srcDieMem;
113 : DeviceMem usroutMem
114 0 : = DeviceMem::create(static_cast<u8*>(opInfo_->outputAddr) + totalDataBytes_ * dieRankId, dataBytes_);
115 : ;
116 :
117 : // 定义本地目标内存和远程源内存
118 0 : if (srcDMAMemSliceId == 0) {
119 0 : locDieDst = usroutMem.range(0, dataPartSize);
120 0 : srcDieMem = DeviceMem::create(static_cast<u8*>(remCCLMemPtr), dataPartSize);
121 : } else {
122 0 : locDieDst = usroutMem.range(dataPartSize, dataBytes_ - dataPartSize);
123 0 : srcDieMem = DeviceMem::create(static_cast<u8*>(remCCLMemPtr) + dataPartSize, dataBytes_ - dataPartSize);
124 : }
125 0 : u32 linkType = static_cast<u32>(links[dieRankId]->GetLinkType());
126 0 : HCCL_DEBUG("[AllGatherHccsSio][RunInterDieOpbase] dstRankId[%u], linkType[%u]", dieRankId, linkType);
127 0 : HCCL_INFO(
128 : "RunInterDieOpbase: dieRankId[%d], locDieDst ptr[%p], locDieDst size[%ld], remCCLMemPtr[%p]", dieRankId,
129 : locDieDst.ptr(), locDieDst.size(), remCCLMemPtr);
130 : // 执行异步内存复制
131 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, locDieDst, srcDieMem, meshStreams_[srcDMAMemSliceId]));
132 :
133 0 : return HCCL_SUCCESS;
134 0 : }
135 :
136 : // allgather的入口函数
137 : HcclResult
138 0 : AllGatherHccsSio::RunAsync(const u32 rank, const u32 rankSize, [[maybe_unused]] const std::vector<LINK>& links)
139 : {
140 : /*rank0:
141 : 从userin拷贝到userout上半部分
142 : userout下半部分的上半部分通过sio读取rank1的userin上半部分
143 : userout下半部分的下半部分通过hccs读取rank1的userin下半部分
144 : */
145 :
146 : /*rank1:
147 : 从userin拷贝到userout下半部分
148 : userout上半部分的上半部分通过sio读取rank0的userin上半部分
149 : userout上半部分的下半部分通过hccs读取rank0的userin下半部分
150 : */
151 0 : intraRankSize_ = rankSize;
152 0 : u32 dieRankId = (rank + 1) % rankSize;
153 : // 数据切分为2
154 : static u32 HCCL_ALLGATHER_SPLIT_FACTOR = 2;
155 :
156 : // dmaMem0部分userin,dmaMem1部分userout
157 0 : DeviceMem dmaMem0 = DeviceMem::create(inputMem_.ptr(), dataBytes_);
158 0 : DeviceMem dmaMem1 = DeviceMem::create(outputMem_.ptr(), dataBytes_ * intraRankSize_);
159 0 : DeviceMem locDieDst = dmaMem1.range(dataBytes_ * rank, dataBytes_);
160 :
161 0 : HCCL_INFO(
162 : "RunAsync: dmaMem0 ptr[%p], dmaMem0 size[%ld]; dmaMem1 ptr[%p], dmaMem1 size[%ld]; locDieDst ptr[%p], "
163 : "locDieDst size[%ld]",
164 : inputMem_.ptr(), dataBytes_, outputMem_.ptr(), dataBytes_ * intraRankSize_, locDieDst.ptr(), dataBytes_);
165 :
166 0 : dmaMem_.push_back(dmaMem0); // userin
167 0 : dmaMem_.push_back(dmaMem1); // userout
168 :
169 0 : if (GetWorkflowMode() == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
170 : // usrin 到 cclin
171 0 : DeviceMem locDieUsrin = DeviceMem::create(static_cast<u8*>(opInfo_->inputAddr), dataBytes_);
172 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dmaMem0, locDieUsrin, stream_));
173 0 : } else {
174 : // step 0操作 : 所有卡本地数据从userIn-->userout
175 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, locDieDst, dmaMem0, stream_));
176 : }
177 :
178 : // 主流启动从流
179 0 : CHK_RET(NotifySubStreamStart());
180 :
181 : // step 1 : die间 && device间并行收发
182 :
183 : // 数据搬运及后同步
184 0 : u32 srcDMAMemSliceId = 0;
185 :
186 0 : CHK_RET(outerCommInfoHccs_.links[dieRankId]->TxAck(meshStreams_[srcDMAMemSliceId])); // AckRecord
187 0 : CHK_RET(outerCommInfoHccs_.links[dieRankId]->RxAck(meshStreams_[srcDMAMemSliceId])); // AckWait
188 0 : CHK_RET(outerCommInfoSio_.links[dieRankId]->TxAck(meshStreams_[srcDMAMemSliceId + 1])); // AckRecord
189 0 : CHK_RET(outerCommInfoSio_.links[dieRankId]->RxAck(meshStreams_[srcDMAMemSliceId + 1])); // AckWait
190 :
191 0 : if (GetWorkflowMode() == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
192 : // 本地userout 读取die间 cclin by sio
193 0 : CHK_RET(RunInterDieOpBase(dieRankId, outerCommInfoHccs_.links, srcDMAMemSliceId));
194 0 : notifyIdx_++;
195 :
196 : // 本地userout 读取die间 userin by hccs
197 : // srcDMAMemSliceId++;
198 0 : CHK_RET(RunInterDieOpBase(dieRankId, outerCommInfoSio_.links, srcDMAMemSliceId + 1));
199 :
200 : // 本地usrout读取本地usrin
201 : DeviceMem locDieSrc = DeviceMem::create(
202 0 : static_cast<u8*>(opInfo_->inputAddr), count_ / HCCL_ALLGATHER_SPLIT_FACTOR * SIZE_TABLE[dataType_]);
203 0 : locDieDst = DeviceMem::create(
204 0 : static_cast<u8*>(opInfo_->outputAddr) + totalDataBytes_ * rank,
205 0 : count_ / HCCL_ALLGATHER_SPLIT_FACTOR * SIZE_TABLE[dataType_]);
206 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, locDieDst, locDieSrc, meshStreams_[srcDMAMemSliceId + 2]));
207 :
208 0 : locDieSrc = DeviceMem::create(
209 0 : static_cast<u8*>(opInfo_->inputAddr) + count_ / HCCL_ALLGATHER_SPLIT_FACTOR * SIZE_TABLE[dataType_],
210 0 : dataBytes_ - count_ / HCCL_ALLGATHER_SPLIT_FACTOR * SIZE_TABLE[dataType_]);
211 0 : locDieDst = DeviceMem::create(
212 0 : static_cast<u8*>(opInfo_->outputAddr) + totalDataBytes_ * rank
213 0 : + count_ / HCCL_ALLGATHER_SPLIT_FACTOR * SIZE_TABLE[dataType_],
214 0 : dataBytes_ - count_ / HCCL_ALLGATHER_SPLIT_FACTOR * SIZE_TABLE[dataType_]);
215 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, locDieDst, locDieSrc, meshStreams_[srcDMAMemSliceId + 3]));
216 0 : } else {
217 : // 本地userout 读取die间 userin by sio
218 0 : CHK_RET(RunInterDie(dieRankId, outerCommInfoHccs_.links, srcDMAMemSliceId));
219 0 : notifyIdx_++;
220 :
221 : // 本地userout 读取die间 userin by hccs
222 : // srcDMAMemSliceId++;
223 0 : CHK_RET(RunInterDie(dieRankId, outerCommInfoSio_.links, srcDMAMemSliceId + 1));
224 : }
225 :
226 0 : CHK_RET(outerCommInfoHccs_.links[dieRankId]->TxDataSignal(meshStreams_[srcDMAMemSliceId])); // DataRecord
227 0 : CHK_RET(outerCommInfoHccs_.links[dieRankId]->RxDataSignal(meshStreams_[srcDMAMemSliceId])); // Datawait
228 0 : CHK_RET(outerCommInfoSio_.links[dieRankId]->TxDataSignal(meshStreams_[srcDMAMemSliceId + 1])); // DataRecord
229 0 : CHK_RET(outerCommInfoSio_.links[dieRankId]->RxDataSignal(meshStreams_[srcDMAMemSliceId + 1])); // Datawait
230 :
231 0 : CHK_RET(WaitSubStreamFinish());
232 :
233 0 : HCCL_INFO("[AllGatherHccsSio][RunAsync]AllGatherHccsSio finished groupRankId[%u] ", userRank_);
234 0 : return HCCL_SUCCESS;
235 0 : }
236 : REGISTER_TEMPLATE(TemplateType::TEMPLATE_ALL_GATHER_HCCS_SIO, AllGatherHccsSio);
237 : } // namespace hccl
|