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_v_2level_pipeline_excecutor.h"
12 : namespace hccl {
13 :
14 0 : CollRunAlltoAllVTwoLevelPipeline::CollRunAlltoAllVTwoLevelPipeline(
15 0 : const HcclDispatcher dispatcher, std::unique_ptr<TopoMatcher>& topoMatcher)
16 0 : : CollAlltoAllExecutor(dispatcher, topoMatcher)
17 0 : {}
18 :
19 : // 计算 alltoall pipeline 910B 的两级流水算法本卡需要的 scratch 大小(图模式需要)
20 0 : u64 CollRunAlltoAllVTwoLevelPipeline::GetAlltoall2LevelPipelineScratchSize910B(
21 : u32 rank, std::vector<SendRecvInfo>& allMeshAggregationSendRecvInfo)
22 : {
23 0 : u64 userRankSize = allMeshAggregationSendRecvInfo.size();
24 0 : u64 maxBlockSize = 0;
25 0 : u64 maxScratchSize = 0;
26 0 : const SendRecvInfo& info = allMeshAggregationSendRecvInfo[rank];
27 0 : for (u64 i = 0; i < userRankSize; i++) {
28 0 : maxBlockSize = std::max(maxBlockSize, info.sendLength[i]);
29 0 : maxBlockSize = std::max(maxBlockSize, info.recvLength[i]);
30 0 : maxScratchSize = std::max(maxScratchSize, info.sendOffset[i] + info.sendLength[i]);
31 0 : maxScratchSize = std::max(maxScratchSize, info.recvOffset[i] + info.recvLength[i]);
32 : }
33 0 : maxScratchSize = std::max(maxBlockSize * userRankSize, maxScratchSize);
34 0 : return maxScratchSize;
35 : }
36 :
37 : // 计算 alltoall pipeline 910B 的两级流水算法所有卡需要的 scratch 大小的最大值(单算子模式需要)
38 0 : u64 CollRunAlltoAllVTwoLevelPipeline::GetAlltoall2LevelPipelineMaxScratchSize910B(
39 : std::vector<SendRecvInfo>& allMeshAggregationSendRecvInfo)
40 : {
41 0 : u64 maxScratchSize = 0;
42 0 : for (u32 rank = 0, userRankSize = allMeshAggregationSendRecvInfo.size(); rank < userRankSize; rank++) {
43 0 : u64 currRankScratchSize = GetAlltoall2LevelPipelineScratchSize910B(rank, allMeshAggregationSendRecvInfo);
44 0 : maxScratchSize = (currRankScratchSize > maxScratchSize ? currRankScratchSize : maxScratchSize);
45 : }
46 0 : return maxScratchSize;
47 : }
48 :
49 0 : HcclResult CollRunAlltoAllVTwoLevelPipeline::CalcScratchMemSize(u64& scratchMemSize)
50 : {
51 0 : scratchMemSize = 0U;
52 0 : u64 tmpMemSize = 0U;
53 0 : if (workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OPS_KERNEL_INFO_LIB
54 0 : && topoAttr_.deviceType == DevType::DEV_TYPE_910B) {
55 : // 图模式才需要申请 scratch 在此只计算scratchMem size
56 0 : tmpMemSize = GetAlltoall2LevelPipelineMaxScratchSize910B(allMeshAggregationSendRecvInfo_);
57 : }
58 0 : scratchMemSize = CalAlltoAllVScratchMemSize(tmpMemSize);
59 0 : HCCL_INFO(
60 : "[CollRunAlltoAllVTwoLevelPipeline][CalcScratchMemSize] tag_[%s] scratchMemSize[%llu]", tag_.c_str(),
61 : scratchMemSize);
62 0 : return HCCL_SUCCESS;
63 : }
64 :
65 0 : HcclResult CollRunAlltoAllVTwoLevelPipeline::CalcStreamNum(u32& streamNum)
66 : {
67 0 : u32 totalStreamNum = topoAttr_.deviceNumPerAggregation + 1U;
68 0 : streamNum = totalStreamNum - 1U;
69 0 : HCCL_INFO("[CollRunAlltoAllVTwoLevelPipeline][CalcStreamNum] tag_[%s] streamNum[%u]", tag_.c_str(), streamNum);
70 0 : return HCCL_SUCCESS;
71 : }
72 :
73 0 : HcclResult CollRunAlltoAllVTwoLevelPipeline::CalcLevel0CommInfo(
74 : TransportMemType inputType, TransportMemType outputType, std::vector<LevelNSubCommTransport>& opTransport)
75 : {
76 0 : CommParaInfo commParaLevel0(COMM_MESH_L0, CommType::COMM_TAG_MESH);
77 0 : CHK_RET(CalcCommPlaneInfo(tag_, commParaLevel0, opTransport[COMM_MESH_L0], inputType, outputType));
78 0 : return HCCL_SUCCESS;
79 0 : }
80 :
81 0 : HcclResult CollRunAlltoAllVTwoLevelPipeline::CalcLevel1CommInfo(
82 : TransportMemType inputType, TransportMemType outputType, std::vector<LevelNSubCommTransport>& opTransport)
83 : {
84 0 : CommParaInfo commParaLevel1(COMM_MESH_L1, CommType::COMM_TAG_MESH);
85 0 : CHK_RET(CalcCommPlaneInfo(tag_, commParaLevel1, opTransport[COMM_MESH_L1], inputType, outputType));
86 0 : return HCCL_SUCCESS;
87 0 : }
88 :
89 0 : HcclResult CollRunAlltoAllVTwoLevelPipeline::CalcCommInfo(std::vector<LevelNSubCommTransport>& opTransport)
90 : {
91 0 : TransportMemType inputType = TransportMemType::RESERVED;
92 0 : TransportMemType outputType = TransportMemType::RESERVED;
93 :
94 0 : CHK_RET(CalNoScratchAlltoallCommInfo(inputType, outputType, opTransport));
95 0 : return HCCL_SUCCESS;
96 : }
97 :
98 0 : HcclResult CollRunAlltoAllVTwoLevelPipeline::CalNoScratchAlltoallCommInfo(
99 : TransportMemType inputType, TransportMemType outputType, std::vector<LevelNSubCommTransport>& opTransport)
100 : {
101 : (void)inputType;
102 : (void)outputType;
103 0 : if (workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE
104 0 : || topoAttr_.deviceType == DevType::DEV_TYPE_910_93) {
105 0 : CHK_RET(CalcLevel0CommInfo(TransportMemType::CCL_OUTPUT, TransportMemType::CCL_OUTPUT, opTransport));
106 0 : CHK_RET(CalcLevel1CommInfo(TransportMemType::CCL_INPUT, TransportMemType::CCL_OUTPUT, opTransport));
107 0 : } else {
108 0 : CHK_RET(CalcLevel0CommInfo(TransportMemType::SCRATCH, TransportMemType::CCL_OUTPUT, opTransport));
109 0 : CHK_RET(CalcLevel1CommInfo(TransportMemType::CCL_INPUT, TransportMemType::SCRATCH, opTransport));
110 : }
111 :
112 0 : return HCCL_SUCCESS;
113 : }
114 :
115 0 : HcclOpMetaInfo CollRunAlltoAllVTwoLevelPipeline::GetOpMeta(HcclCMDType opType, const u64 size)
116 : {
117 : (void)opType;
118 : (void)size;
119 0 : bool hugeData = (isAlltoAllZCopyMode_) ? (algResResp_->paramInputMem.size() > SDMA_SEND_MAX_SIZE) : (false);
120 : bool alltoallPingPong
121 0 : = ((workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE
122 0 : || topoAttr_.deviceType == DevType::DEV_TYPE_910_93)
123 0 : && !static_cast<bool>(topoAttr_.multiModuleDiffDeviceNumMode)
124 0 : && GetAlltoall2LevelPipelineMaxScratchSize910B(allMeshAggregationSendRecvInfo_)
125 0 : > algResResp_->cclInputMem.size());
126 0 : HcclOpMetaInfo opMeta;
127 0 : if (AlltoAllVParam_.opType == HcclCMDType::HCCL_CMD_ALLTOALLV) {
128 0 : opMeta = HcclOpMetaInfo::GetOneForAllToAllV(
129 0 : (isAlltoAllZCopyMode_ ? CopyPattern::ZCOPY : CopyPattern::BCOPY), algResResp_->paramInputMem.size(),
130 0 : hugeData || alltoallPingPong);
131 : } else {
132 0 : opMeta = HcclOpMetaInfo::GetOneForAllToAllV(
133 0 : (isAlltoAllZCopyMode_ ? CopyPattern::ZCOPY : CopyPattern::BCOPY), algResResp_->paramInputMem.size(),
134 0 : hugeData || alltoallPingPong);
135 : }
136 0 : HCCL_DEBUG("[CollRunAlltoAllVTwoLevelPipeline][GetOpMeta] Get OpMeta for AllToAllV pipeline success.");
137 0 : return opMeta;
138 : }
139 :
140 0 : HcclResult CollRunAlltoAllVTwoLevelPipeline::KernelRun(const OpParam& param, ExecMem& execMem)
141 : {
142 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[CollRunAlltoAllVTwoLevelPipeline][KernelRun] AllToAllV two level pipeline start");
143 :
144 0 : bool cclEnough = true;
145 0 : if ((workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE
146 0 : || topoAttr_.deviceType == DevType::DEV_TYPE_910_93)
147 0 : && GetAlltoall2LevelPipelineMaxScratchSize910B(allMeshAggregationSendRecvInfo_) > execMem.inputMem.size()) {
148 0 : cclEnough = false;
149 : }
150 0 : HCCL_CONFIG_INFO(
151 : HCCL_ALG, "[CollRunAlltoAllVTwoLevelPipeline][KernelRun] AllToAllV pipeline run %s algo",
152 : cclEnough ? "cclEnough" : "ping pong");
153 0 : A2aPipelineMemory a2aPipelineMemory;
154 0 : a2aPipelineMemory.userInput = algResResp_->paramInputMem;
155 0 : a2aPipelineMemory.userOutput = algResResp_->paramOutputMem;
156 : // 具体传入 A2aPipelineMemory 对象的 alltoall pipeline executor 会根据图模式还是单算子模式
157 : // 选择使用 ccl 还是 scratch,不会访问空指针
158 0 : if (workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE
159 0 : || topoAttr_.deviceType == DevType::DEV_TYPE_910_93) {
160 0 : a2aPipelineMemory.cclInBuffer = execMem.inputMem;
161 0 : a2aPipelineMemory.cclOutBuffer = execMem.outputMem;
162 : } else {
163 0 : a2aPipelineMemory.scratchMem = execMem.scratchMem;
164 : }
165 0 : std::unique_ptr<AlgTemplateBase> alltoallPipe = nullptr;
166 0 : if (cclEnough) {
167 0 : alltoallPipe = AlgTemplateRegistry::Instance().GetAlgTemplate(
168 0 : TemplateType::TEMPLATE_ALL_2_ALL_PIPELINE_MESH_PAIRWISE_CCL_ENOUGH, dispatcher_);
169 : } else {
170 0 : alltoallPipe = AlgTemplateRegistry::Instance().GetAlgTemplate(
171 0 : TemplateType::TEMPLATE_ALL_2_ALL_PIPELINE_MESH_PAIRWISE_PING_PONG, dispatcher_);
172 : }
173 :
174 0 : CHK_RET(CheckCommSize(COMM_MESH_L0, COMM_INDEX_0 + 1));
175 0 : SubCommInfo level0CommInfo = GetSubCommInfo(COMM_MESH_L0, COMM_INDEX_0);
176 0 : CHK_RET(CheckCommSize(COMM_MESH_L1, COMM_INDEX_0 + 1));
177 0 : SubCommInfo level1CommInfo = GetSubCommInfo(COMM_MESH_L1, COMM_INDEX_0);
178 :
179 0 : CHK_SMART_PTR_NULL(alltoallPipe);
180 0 : CHK_RET(alltoallPipe->Prepare(
181 : topoAttr_.userRank, a2aPipelineMemory, level0CommInfo, level1CommInfo, const_cast<Stream&>(param.stream),
182 : algResResp_->slaveStreams, algResResp_->notifiesMain, algResResp_->notifiesAux, allMeshAggregationSendRecvInfo_,
183 : workflowMode_));
184 0 : CHK_RET(alltoallPipe->RunAsync());
185 0 : HCCL_INFO("[CollRunAlltoAllVTwoLevelPipeline][kernelRun] AllToAllV two level pipeline exec end");
186 0 : return HCCL_SUCCESS;
187 0 : }
188 :
189 : REGISTER_EXEC("RunAlltoAllVTwoLevelPipeline", AlltoAllVTwoLevelPipeline, CollRunAlltoAllVTwoLevelPipeline);
190 : } // namespace hccl
|