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