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_v_operator.h"
12 : #include "device_capacity.h"
13 : #include "coll_alg_op_registry.h"
14 : #include "hccl_aiv.h"
15 :
16 : namespace hccl {
17 :
18 : constexpr u64 MAX_310P_RANK_SIZE = 4;
19 : constexpr u32 MODULE_NUM_FOUR = 4;
20 :
21 0 : AllGatherVOperator::AllGatherVOperator(
22 : AlgConfigurator* algConfigurator, CCLBufferManager& cclBufferManager, HcclDispatcher dispatcher,
23 0 : std::unique_ptr<TopoMatcher>& topoMatcher)
24 0 : : CollAlgOperator(algConfigurator, cclBufferManager, dispatcher, topoMatcher, HcclCMDType::HCCL_CMD_ALLGATHER_V)
25 0 : {}
26 :
27 0 : AllGatherVOperator::~AllGatherVOperator() {}
28 :
29 : HcclResult
30 0 : AllGatherVOperator::SelectAlg(const std::string& tag, const OpParam& param, std::string& algName, std::string& newTag)
31 : {
32 : HcclResult ret;
33 0 : HCCL_DEBUG("[%s] SelectAlg begins", __func__);
34 0 : if (isDiffDeviceType_) {
35 0 : HCCL_ERROR("[AllGatherVOperator][SelectAlg] AllGatherV not support diffDeviceType");
36 0 : return HCCL_E_NOT_SUPPORT;
37 0 : } else if (deviceType_ == DevType::DEV_TYPE_910_93) {
38 0 : ret = SelectAlgfor91093(param, algName);
39 0 : } else if (deviceType_ == DevType::DEV_TYPE_910B) {
40 0 : ret = SelectAlgfor910B(param, algName);
41 0 : } else if (deviceType_ == DevType::DEV_TYPE_310P3) {
42 0 : ret = SelectAlgfor310P3(param, algName);
43 : } else {
44 0 : HCCL_ERROR("[AllGatherVOperator][SelectAlg] AllGatherV only support A3, A2 and 310P.");
45 0 : return HCCL_E_NOT_SUPPORT;
46 : }
47 0 : CHK_PRT_RET(
48 : ret != HCCL_SUCCESS,
49 : HCCL_ERROR("[AllGatherVOperator][SelectAlg]tag[%s], AllGatherV failed, return[%d]", tag.c_str(), ret), ret);
50 :
51 0 : if (workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OPS_KERNEL_INFO_LIB) {
52 0 : newTag = tag;
53 0 : } else if (deviceType_ == DevType::DEV_TYPE_310P3) {
54 0 : newTag = tag + algName;
55 : } else {
56 0 : AlgTypeLevel1 algType1 = algType_.algoLevel1;
57 0 : auto level1Iter = HCCL_ALGO_LEVEL1_NAME_MAP.find(algType1);
58 0 : CHK_PRT_RET(
59 : level1Iter == HCCL_ALGO_LEVEL1_NAME_MAP.end(), HCCL_ERROR("level1: algType1[%u] is invalid.", algType1),
60 : HCCL_E_INTERNAL);
61 :
62 0 : newTag = tag + level1Iter->second + algName;
63 : }
64 :
65 0 : newTag += (param.aicpuUnfoldMode ? "_device" : "_host");
66 0 : return ret;
67 : }
68 :
69 0 : HcclResult AllGatherVOperator::SelectAlgfor91093(const OpParam& param, std::string& algName)
70 : {
71 0 : const HcclDataType dataType = param.VDataDes.dataType;
72 0 : const auto* countsPtr = static_cast<const u64*>(param.VDataDes.counts);
73 0 : const auto countsPerRank = std::vector<u64>(countsPtr, countsPtr + userRankSize_);
74 0 : const u64 maxCount = *std::max_element(countsPerRank.begin(), countsPerRank.end());
75 0 : const u32 unitSize = SIZE_TABLE[dataType];
76 0 : const u64 dataSize = maxCount * unitSize; // 单位:字节
77 0 : if (dataSize >= cclBufferManager_.GetInCCLbufferSize()) {
78 0 : HCCL_WARNING(
79 : "The current inCCLbufferSize is [%llu] bytes, change the HCCL_BUFFSIZE environment variable to "
80 : "be greater than the current data volume[%llu] bytes to improve the performance of the 91093 environment.",
81 : cclBufferManager_.GetInCCLbufferSize(), dataSize);
82 : }
83 :
84 0 : if (multiModuleDiffDeviceNumMode_ || multiSuperPodDiffServerNumMode_) {
85 0 : HCCL_ERROR(
86 : "[AllGatherVOperator][SelectAlgfor91093]not support mode, multiModuleDiffDeviceNumMode_[%u], "
87 : "multiSuperPodDiffServerNumMode_[%u]",
88 : multiModuleDiffDeviceNumMode_, multiSuperPodDiffServerNumMode_);
89 0 : return HCCL_E_NOT_SUPPORT;
90 : } else {
91 0 : if (!(algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_RING
92 0 : || algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_WHOLE_RING
93 0 : || algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NB)) {
94 0 : algType_.algoLevel1 = AlgTypeLevel1::ALG_LEVEL1_NHR;
95 0 : HCCL_WARNING("[AllGatherVOperator][SelectAlgfor91093] only support ring, NB and NHR in AlgoLevel1 yet, "
96 : "default is algType=NHR.");
97 : }
98 0 : if (topoType_ == TopoType::TOPO_TYPE_NP_DOUBLE_RING) {
99 0 : algName = "AlignedAllGatherVDoubleRingFor91093Executor";
100 : } else {
101 0 : algName = "AllGatherVRingFor91093Executor";
102 : }
103 : }
104 :
105 0 : HCCL_INFO("[SelectAlgfor91093] AllGatherV SelectAlgfor91093 is algName [%s]", algName.c_str());
106 0 : return HCCL_SUCCESS;
107 0 : }
108 :
109 0 : HcclResult AllGatherVOperator::SelectAlgfor910B(const OpParam& param, std::string& algName)
110 : {
111 0 : const auto* countsPtr = static_cast<const u64*>(param.VDataDes.counts);
112 0 : auto countsPerRank = std::vector<u64>(countsPtr, countsPtr + userRankSize_);
113 0 : u64 maxCount = *std::max_element(countsPerRank.begin(), countsPerRank.end());
114 0 : u32 unitSize = SIZE_TABLE[param.VDataDes.dataType];
115 0 : u64 dataSize = maxCount * unitSize;
116 0 : bool isBigData = false;
117 :
118 0 : if (dataSize > AIV_ALL_GATHER_SMALL_SIZE) {
119 0 : isBigData = true;
120 : }
121 :
122 0 : if (workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE && !isSingleMeshAggregation_) {
123 0 : u64 cclBufferSize = cclBufferManager_.GetOutCCLbufferSize() / userRankSize_;
124 0 : std::string algTypeLevel1Tag;
125 0 : CHK_RET(AutoSelectAlgTypeLevel1(HcclCMDType::HCCL_CMD_ALLGATHER_V, dataSize, cclBufferSize, algTypeLevel1Tag));
126 0 : if (GetExternalInputHcclEnableEntryLog() && param.opBaseAtraceInfo != nullptr) {
127 0 : CHK_RET(param.opBaseAtraceInfo->SavealgtypeTraceInfo(algTypeLevel1Tag, param.tag));
128 : }
129 0 : }
130 :
131 : // pipeline算法task数量多,如果超出FFTS子图限制,则重定向到NHR算法
132 0 : if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_PIPELINE) {
133 0 : u32 contextNum = CalcContextNumForPipeline(HcclCMDType::HCCL_CMD_ALLGATHER_V);
134 0 : if (contextNum > HCCL_FFTS_CAPACITY) {
135 0 : algType_.algoLevel1 = AlgTypeLevel1::ALG_LEVEL1_NHR;
136 0 : HCCL_WARNING(
137 : "[AllGatherVOperator][SelectAlgfor910B] context num[%u] is out of capacity of FFTS+ graph[%u], "
138 : "reset algorithm to NHR.",
139 : contextNum, HCCL_FFTS_CAPACITY);
140 : }
141 : }
142 :
143 0 : bool isAivMode = topoMatcher_->GetAivModeConfig() && isSingleMeshAggregation_
144 0 : && IsSupportAIVCopy(param.VDataDes.dataType) && dataSize <= AIV_BIG_SIZE;
145 :
146 0 : if (GetWorkflowMode() == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
147 0 : if (isAivMode) {
148 0 : algName = isBigData ? "AllGatherVMeshAivExecutor" : "AllGatherVMeshAivSmallCountExecutor";
149 0 : } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_PIPELINE && !isSingleMeshAggregation_) {
150 0 : algName = "AllGatherVMeshOpbasePipelineExecutor";
151 : } else {
152 0 : algName = "AllGatherVMeshExecutor";
153 : }
154 0 : } else if (GetWorkflowMode() == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OPS_KERNEL_INFO_LIB) {
155 0 : if (isSingleMeshAggregation_) {
156 0 : algName = "AllGatherVMeshGraphExecutor";
157 0 : } else if (
158 0 : deviceNumPerAggregation_ > 1
159 0 : && (dataSize > HCCL_SMALL_COUNT_1_MB || moduleNum_ <= MODULE_NUM_FOUR
160 0 : || algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_PIPELINE)) {
161 0 : algName = "AllGatherVMeshGraphPipelineExecutor";
162 : } else {
163 0 : algName = "AllGatherVMeshExecutor";
164 : }
165 : }
166 0 : HCCL_INFO("[SelectAlgfor910B] AllGatherV SelectAlgfor910B is algName [%s]", algName.c_str());
167 0 : return HCCL_SUCCESS;
168 0 : }
169 :
170 0 : HcclResult AllGatherVOperator::SelectAlgfor310P3(const OpParam& param, std::string& algName)
171 : {
172 : (void)param;
173 0 : CHK_PRT_RET(
174 : userRankSize_ > MAX_310P_RANK_SIZE,
175 : HCCL_ERROR(
176 : "[AllGatherVOperator][SelectAlgfor310P3]rankSize[%u] is not supported.AllGatherV does not support the "
177 : "scenario where the rankSize is greater than 4.",
178 : userRankSize_),
179 : HCCL_E_NOT_SUPPORT);
180 0 : algName = "AllGatherVFor310PExecutor";
181 0 : HCCL_INFO("[SelectAlgfor310P3] AllGatherV SelectAlgfor310P3 is algName [%s]", algName.c_str());
182 0 : return HCCL_SUCCESS;
183 : }
184 :
185 : REGISTER_OP(HcclCMDType::HCCL_CMD_ALLGATHER_V, AllGatherV, AllGatherVOperator);
186 :
187 : } // namespace hccl
|