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 "coll_all_to_all_symmetric_memory_executor.h"
12 :
13 : namespace hccl {
14 :
15 0 : CollRunAlltoAllFullMeshSymmetricMemory::CollRunAlltoAllFullMeshSymmetricMemory(const HcclDispatcher dispatcher,
16 0 : std::unique_ptr<TopoMatcher> &topoMatcher)
17 0 : : CollAlltoAllExecutor(dispatcher, topoMatcher)
18 : {
19 0 : desc_.isZeroCopy = true;
20 0 : }
21 :
22 0 : HcclResult CollRunAlltoAllFullMeshSymmetricMemory::Orchestrate(OpParam& param, AlgResourceResponse& algRes)
23 : {
24 0 : HcclUs startut = TIME_NOW();
25 0 : HcclResult ret = HCCL_SUCCESS;
26 0 : tag_ = param.tag;
27 0 : algResResp_ = &algRes;
28 0 : AlltoAllVParam_ = param;
29 :
30 0 : ExecMem execMem;
31 0 : execMem.count = 0;
32 0 : execMem.inputPtr = param.inputPtr;
33 0 : execMem.outputPtr = param.outputPtr;
34 0 : execMem.inputMem = algRes.paramInputMem;
35 0 : execMem.outputMem = algRes.paramOutputMem;
36 0 : ret = KernelRun(param, execMem);
37 :
38 0 : CHK_PRT_RET(ret != HCCL_SUCCESS,
39 : HCCL_ERROR("[CollRunAlltoAllFullMeshSymmetricMemory][Orchestrate]errNo[0x%016llx]executor run failed",
40 : HCCL_ERROR_CODE(ret)), ret);
41 :
42 0 : HCCL_INFO("tag[%s], AlltoAllFullMeshSymmetricMemory tempAlg orchestrate success, take time [%lld]us.",
43 : param.tag.c_str(), DURATION_US(TIME_NOW() - startut));
44 0 : return HCCL_SUCCESS;
45 0 : }
46 :
47 0 : HcclResult CollRunAlltoAllFullMeshSymmetricMemory::GetLocalSDMAGroupInfo(u32& devNumInlocalPod, u32& rankIdxInPod) const
48 : {
49 0 : CHK_RET(topoMatcher_->GetLocalSuperPodRankSize(topoAttr_.userRank, devNumInlocalPod, rankIdxInPod));
50 0 : CHK_PRT_RET(devNumInlocalPod == INVALID_VALUE_RANKSIZE,
51 : HCCL_ERROR("[CollRunAlltoAllFullMeshSymmetricMemory][GetLocalSDMAGroupInfo]get local superPod total ranksize failed."),
52 : HCCL_E_PARA);
53 0 : return HCCL_SUCCESS;
54 : }
55 :
56 0 : HcclResult CollRunAlltoAllFullMeshSymmetricMemory::CalcStreamNum(u32& streamNum)
57 : {
58 : // 每个超节点内的卡数
59 0 : u32 devNumInlocalPod = INVALID_VALUE_RANKSIZE;
60 0 : u32 rankIdxInPod = INVALID_VALUE_RANKID;
61 0 : CHK_RET(GetLocalSDMAGroupInfo(devNumInlocalPod, rankIdxInPod));
62 :
63 : // 单超节点场景需要的从流数量,待确认是否需要减去一条主流
64 0 : streamNum = (devNumInlocalPod > ALLTOALLV_DIRECT_FULLMESH_SDMA_CONCURRENT_SIZE) ?
65 0 : (ALLTOALLV_DIRECT_FULLMESH_SDMA_CONCURRENT_SIZE) : (devNumInlocalPod);
66 :
67 0 : HCCL_INFO("[CollRunAlltoAllFullMeshSymmetricMemory][CalcStreamNum] tag[%s] streamNum[%u]",
68 : tag_.c_str(), streamNum);
69 0 : return HCCL_SUCCESS;
70 : }
71 :
72 : // level0-level1 打平fullmesh
73 : // 超节点内建SDMA链路;超节点间建RDMA链路
74 0 : HcclResult CollRunAlltoAllFullMeshSymmetricMemory::CalcLevel0CommInfo(TransportMemType inputType, TransportMemType outputType,
75 : std::vector<LevelNSubCommTransport>& opTransport)
76 : {
77 0 : CommParaInfo commCombinePara(COMM_COMBINE_ORDER, CommType::COMM_TAG_MESH);
78 0 : CHK_RET(CalcCommPlaneInfo(tag_, commCombinePara, opTransport[COMM_COMBINE_ORDER], inputType, outputType));
79 0 : LevelNSubCommTransport &commTransportLevel0 = opTransport[COMM_COMBINE_ORDER];
80 0 : for (u32 subCommIndex = 0; subCommIndex < commTransportLevel0.size(); subCommIndex++) {
81 0 : commTransportLevel0[subCommIndex].isZeroCopy = true;
82 : }
83 0 : return HCCL_SUCCESS;
84 0 : }
85 :
86 0 : HcclResult CollRunAlltoAllFullMeshSymmetricMemory::CalcTransportMemType(TransportMemType &inputType, TransportMemType &outputType) const
87 : {
88 0 : inputType = TransportMemType::CCL_INPUT;
89 0 : outputType = TransportMemType::CCL_OUTPUT;
90 :
91 0 : HCCL_INFO("[CollRunAlltoAllFullMeshSymmetricMemory][CalcTransportMemType] tag[%s] inputType[%d], outputType[%d]",
92 : tag_.c_str(), inputType, outputType);
93 0 : return HCCL_SUCCESS;
94 : }
95 :
96 0 : HcclResult CollRunAlltoAllFullMeshSymmetricMemory::CalcCommInfo(std::vector<LevelNSubCommTransport>& opTransport)
97 : {
98 0 : TransportMemType inputType = TransportMemType::RESERVED;
99 0 : TransportMemType outputType = TransportMemType::RESERVED;
100 :
101 0 : CHK_RET(CalcTransportMemType(inputType, outputType));
102 : // level0 - level1 全连接通信域
103 0 : CHK_RET(CalcLevel0CommInfo(inputType, outputType, opTransport));
104 0 : return HCCL_SUCCESS;
105 : }
106 :
107 0 : HcclResult CollRunAlltoAllFullMeshSymmetricMemory::GetLocalSendRecvInfoforAlltoall(const OpParam ¶m)
108 : {
109 0 : u64 curRecvOffset = 0;
110 0 : for (u32 j = 0; j < topoAttr_.userRankSize; j++) {
111 0 : u64 curSendCounts = param.All2AllDataDes.sendCount;
112 0 : u64 curSendLength = curSendCounts * SIZE_TABLE[param.All2AllDataDes.sendType];
113 0 : sendRecvInfo_.remoteSendOffset[j] = curSendLength * topoAttr_.userRank;
114 :
115 0 : u64 curRecvCounts = param.All2AllDataDes.sendCount;
116 0 : u64 curRecvLength = curRecvCounts * SIZE_TABLE[param.All2AllDataDes.recvType];
117 0 : sendRecvInfo_.localRecvLength[j] = curRecvLength;
118 0 : sendRecvInfo_.localRecvOffset[j] = curRecvOffset;
119 0 : curRecvOffset += curRecvLength;
120 0 : HCCL_DEBUG("GetLocalSendRecvInfoforAlltoall rank[%u], remoteSendOffset[j][%llu], localRecvLength[j][%llu] "\
121 : "localRecvOffset[j][%llu]", topoAttr_.userRank, sendRecvInfo_.remoteSendOffset[j],
122 : sendRecvInfo_.localRecvLength[j], sendRecvInfo_.localRecvOffset[j]);
123 : }
124 0 : return HCCL_SUCCESS;
125 : }
126 :
127 0 : HcclResult CollRunAlltoAllFullMeshSymmetricMemory::GetLocalSendRecvInfoforAlltoallVC(const OpParam ¶m)
128 : {
129 0 : u64 rankSize = topoAttr_.userRankSize;
130 0 : u64 usrRank = topoAttr_.userRank;
131 0 : for (u32 j = 0; j < topoAttr_.userRankSize; j++) {
132 0 : sendRecvInfo_.remoteSendOffset[j] = 0;
133 0 : for (u32 i = 0; i < usrRank; i++) {
134 0 : u64 sendCounts = *(static_cast<const u64 *>(param.All2AllDataDes.sendCountMatrix) + i + rankSize * j);
135 0 : u64 sendLength = sendCounts * SIZE_TABLE[param.All2AllDataDes.recvType];
136 0 : sendRecvInfo_.remoteSendOffset[j] += sendLength;
137 : }
138 0 : sendRecvInfo_.localRecvOffset[j] = 0;
139 0 : for (u32 i = 0; i < j; i++) {
140 0 : u64 recvCounts = *(static_cast<const u64 *>(param.All2AllDataDes.sendCountMatrix) + usrRank + rankSize * i);
141 0 : u64 recvLength = recvCounts * SIZE_TABLE[param.All2AllDataDes.sendType];
142 0 : sendRecvInfo_.localRecvOffset[j] += recvLength;
143 : }
144 0 : u64 curRecvCounts = *(static_cast<const u64 *>(param.All2AllDataDes.sendCountMatrix) + usrRank + rankSize * j);
145 0 : sendRecvInfo_.localRecvLength[j] = curRecvCounts * SIZE_TABLE[param.All2AllDataDes.recvType];
146 :
147 0 : HCCL_DEBUG("GetLocalSendRecvInfoforAlltoallVC rank[%u], remoteSendOffset[%llu], "\
148 : "localRecvLength[%llu], localRecvOffset[%llu]", topoAttr_.userRank, sendRecvInfo_.remoteSendOffset[j],
149 : sendRecvInfo_.localRecvLength[j], sendRecvInfo_.localRecvOffset[j]);
150 : }
151 0 : return HCCL_SUCCESS;
152 : }
153 :
154 0 : HcclResult CollRunAlltoAllFullMeshSymmetricMemory::GetAlltoAllTmpRankSendRecvInfo(const OpParam ¶m)
155 : {
156 0 : sendRecvInfo_.remoteSendOffset.resize(topoAttr_.userRankSize, 0);
157 0 : sendRecvInfo_.localRecvLength.resize(topoAttr_.userRankSize, 0);
158 0 : sendRecvInfo_.localRecvOffset.resize(topoAttr_.userRankSize, 0);
159 :
160 0 : if (param.opType == HcclCMDType::HCCL_CMD_ALLTOALL) {
161 0 : CHK_RET(GetLocalSendRecvInfoforAlltoall(param));
162 0 : } else if (param.opType == HcclCMDType::HCCL_CMD_ALLTOALLVC) {
163 0 : CHK_RET(GetLocalSendRecvInfoforAlltoallVC(param));
164 : } else {
165 0 : HCCL_ERROR("Only support optype AllToAll and AllToAllVC !");
166 : }
167 0 : return HCCL_SUCCESS;
168 : }
169 :
170 0 : HcclResult CollRunAlltoAllFullMeshSymmetricMemory::KernelRun(const OpParam ¶m, ExecMem &execMem)
171 : {
172 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] AllToAll fullmesh start.", __func__);
173 :
174 : // 准备数据
175 0 : CHK_RET(ActiveSlaveStreams(param.stream));
176 0 : CHK_RET(GetAlltoAllTmpRankSendRecvInfo(param));
177 :
178 : // 获取当前超节点内总卡数
179 0 : u32 devNumInlocalPod = INVALID_VALUE_RANKSIZE;
180 0 : u32 rankIdxInPod = INVALID_VALUE_RANKID;
181 0 : CHK_RET(GetLocalSDMAGroupInfo(devNumInlocalPod, rankIdxInPod));
182 :
183 : // 获取通信域
184 0 : CHK_RET(CheckCommSize(COMM_COMBINE_ORDER, COMM_INDEX_0 + 1));
185 0 : SubCommInfo level0CommInfo = GetSubCommInfo(COMM_COMBINE_ORDER, COMM_INDEX_0);
186 : // isSuPodAsym 表示A2A3卡数不一致场景或者A3多超节点server数不同场景
187 0 : bool isSuPodAsym = false;
188 :
189 : // 执行
190 0 : std::unique_ptr<AlgTemplateBase> tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
191 0 : TemplateType::TEMPLATE_ALL_2_ALL_FULL_MESH_SYMMETRIC_MEMORY, dispatcher_);
192 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_2_ALL_FULL_MESH_SYMMETRIC_MEMORY in COMM_COMBINE_ORDER", __func__);
193 0 : CHK_SMART_PTR_NULL(tempAlg);
194 :
195 0 : PrepareData prepareData;
196 0 : prepareData.stream = param.stream;
197 0 : prepareData.userRank = topoAttr_.userRank;
198 0 : prepareData.userRankSize = topoAttr_.userRankSize;
199 0 : prepareData.linksPtr = &level0CommInfo.links;
200 0 : prepareData.sendRecvInfoPtr = &sendRecvInfo_;
201 0 : prepareData.devNumInlocalPod = devNumInlocalPod;
202 0 : prepareData.rankIdxInPod = rankIdxInPod;
203 :
204 0 : prepareData.inputMem = algResResp_->paramInputMem;
205 0 : prepareData.outputMem = algResResp_->paramOutputMem;
206 0 : prepareData.cclInMem = execMem.inputMem;
207 0 : prepareData.cclOutMem = execMem.outputMem;
208 0 : prepareData.workMode = workflowMode_;
209 0 : prepareData.subStreamsPtr = &algResResp_->slaveStreams;
210 0 : prepareData.signalPtr = &algResResp_->notifiesMain;
211 0 : prepareData.signalAuxPtr = &algResResp_->notifiesAux;
212 0 : prepareData.isSuPodAsym = isSuPodAsym;
213 0 : prepareData.opType = param.opType;
214 0 : prepareData.algOpContext = algOpContext_;
215 :
216 0 : CHK_RET(tempAlg->Prepare(prepareData));
217 :
218 0 : CHK_RET(tempAlg->RunAsync());
219 :
220 0 : HCCL_INFO("[CollRunAlltoAllFullMeshSymmetricMemory] executor run success.");
221 0 : if (algOpContext_.opRetryHandler.isPostSync == true) {
222 0 : OpParam postSyncParam = param;
223 0 : if ((*prepareData.subStreamsPtr).size() == 0) {
224 0 : CHK_RET(PostSyncWithoutSubstream(postSyncParam, execMem));
225 : } else {
226 0 : PrepareData postSyncPrepareData = prepareData;
227 0 : CHK_RET(PostSyncWithSubstream(postSyncParam, execMem, postSyncPrepareData));
228 0 : }
229 0 : }
230 0 : return HCCL_SUCCESS;
231 0 : }
232 :
233 : REGISTER_EXEC("RunAlltoAllFullMeshSymmetricMemory", AlltoAllFullMeshSymmetricMemory, CollRunAlltoAllFullMeshSymmetricMemory);
234 : } // namespace hccl
|