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_broadcast_comm_executor.h"
12 :
13 : namespace hccl {
14 :
15 1 : CollBroadcastCommExecutor::CollBroadcastCommExecutor(const HcclDispatcher dispatcher,
16 1 : std::unique_ptr<TopoMatcher> &topoMatcher)
17 1 : : CollBroadcastExecutor(dispatcher, topoMatcher)
18 : {
19 1 : }
20 :
21 :
22 0 : HcclResult CollBroadcastCommExecutor::CalcCommInfo(std::vector<LevelNSubCommTransport>& opTransport)
23 : {
24 0 : TransportMemType inputType = TransportMemType::RESERVED;
25 0 : TransportMemType outputType = TransportMemType::RESERVED;
26 0 : CHK_RET(CalcTransportMemType(inputType, outputType));
27 0 : CHK_RET(CalcCombinedCommInfo(inputType, outputType, opTransport));
28 0 : return HCCL_SUCCESS;
29 : }
30 :
31 0 : HcclResult CollBroadcastCommExecutor::CalcCombinedCommInfo(TransportMemType inputType,
32 : TransportMemType outputType,
33 : std::vector<LevelNSubCommTransport>& opTransport)
34 : {
35 0 : CommPlane commPlane = COMM_COMBINE;
36 0 : if (topoAttr_.deviceType == DevType::DEV_TYPE_910_93) {
37 0 : commPlane = COMM_COMBINE_ORDER;
38 : }
39 :
40 0 : CommParaInfo commParaInfo(commPlane, CommType::COMM_TAG_MAX);
41 0 : if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NHR) {
42 0 : commParaInfo.commType = CommType::COMM_TAG_NONUNIFORM_HIERARCHICAL_RING;
43 0 : } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NHR_V1) {
44 0 : commParaInfo.commType = CommType::COMM_TAG_NONUNIFORM_HIERARCHICAL_RING_V1;
45 0 : } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NB) {
46 0 : commParaInfo.commType = CommType::COMM_TAG_NONUNIFORM_BRUCK;
47 : } else {
48 0 : commParaInfo.commType = CommType::COMM_TAG_RING_INNER;
49 : }
50 0 : CHK_RET(CalcCommPlaneInfo(tag_, commParaInfo, opTransport[commPlane], inputType, outputType));
51 :
52 0 : return HCCL_SUCCESS;
53 0 : }
54 :
55 0 : HcclResult CollBroadcastCommExecutor::CalcStreamNum(u32& streamNum)
56 : {
57 : // 只传递从流数量
58 0 : streamNum = 0;
59 0 : HCCL_INFO("[CollBroadcastCommExecutor][CalcStreamNum]tag[%s] streamNum_ is [%u]", tag_.c_str(), streamNum);
60 0 : return HCCL_SUCCESS;
61 : }
62 :
63 0 : void SetPrepareData(PrepareData &prepareData, const OpParam ¶m,
64 : const ExecMem &execMem, const u32 &rootRank)
65 : {
66 0 : prepareData.inputMem = execMem.inputMem;
67 0 : prepareData.outputMem = execMem.outputMem;
68 0 : prepareData.scratchMem = execMem.outputMem;
69 0 : prepareData.count = execMem.count;
70 0 : prepareData.dataType = param.DataDes.dataType;
71 0 : prepareData.stream = param.stream;
72 0 : prepareData.reductionOp = HCCL_REDUCE_RESERVED;
73 0 : prepareData.root = rootRank;
74 0 : prepareData.baseOffset = 0;
75 0 : }
76 :
77 0 : HcclResult CollBroadcastCommExecutor::KernelRun(const OpParam ¶m, ExecMem &execMem)
78 : {
79 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[CollBroadcastCommExecutor][KernelRun] userRank[%u] starts.", topoAttr_.userRank);
80 0 : CommPlane commPlane = COMM_COMBINE;
81 0 : if (topoAttr_.deviceType == DevType::DEV_TYPE_910_93) {
82 0 : commPlane = COMM_COMBINE_ORDER;
83 : }
84 :
85 0 : CHK_RET(CheckCommSize(commPlane, COMM_INDEX_0 + 1));
86 0 : SubCommInfo combinedCommInfo = GetSubCommInfo(commPlane, COMM_INDEX_0);
87 :
88 0 : bool isUsedRegister = false;
89 0 : std::unique_ptr<AlgTemplateBase> tempAlg;
90 0 : u64 curSize = execMem.count * SIZE_TABLE[param.DataDes.dataType];
91 0 : if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NHR) {
92 0 : if (curSize <= NHR_BCAST_SMALL_SIZE) {
93 0 : tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
94 0 : TemplateType::TEMPLATE_BROADCAST_NHR_ONESHOT, dispatcher_);
95 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_BROADCAST_NHR_ONESHOT in COMM_COMBINE/COMM_COMBINE_ORDER", __func__);
96 : } else {
97 0 : tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
98 0 : TemplateType::TEMPLATE_BROADCAST_NHR, dispatcher_);
99 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_BROADCAST_NHR in COMM_COMBINE/COMM_COMBINE_ORDER", __func__);
100 : }
101 0 : } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NHR_V1) {
102 0 : isUsedRegister = true;
103 0 : tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_BROADCAST_NHR_V1,
104 0 : dispatcher_);
105 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_BROADCAST_NHR_V1 in COMM_COMBINE/COMM_COMBINE_ORDER", __func__);
106 0 : } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NB) {
107 0 : if (ShouldUseBinaryBroadcastOfNB(curSize, combinedCommInfo.localRankSize, topoAttr_.userRankSize,
108 0 : topoAttr_.deviceNumPerAggregation)) {
109 0 : tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
110 0 : TemplateType::TEMPLATE_BROADCAST_NB_BINARY, dispatcher_);
111 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_BROADCAST_NB_BINARY in COMM_COMBINE/COMM_COMBINE_ORDER", __func__);
112 : } else {
113 0 : tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
114 0 : TemplateType::TEMPLATE_BROADCAST_NB, dispatcher_);
115 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_BROADCAST_NB in COMM_COMBINE/COMM_COMBINE_ORDER", __func__);
116 : }
117 : } else {
118 0 : tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
119 0 : TemplateType::TEMPLATE_BROADCAST_RING, dispatcher_);
120 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_BROADCAST_RING in COMM_COMBINE/COMM_COMBINE_ORDER", __func__);
121 : }
122 0 : CHK_SMART_PTR_NULL(tempAlg);
123 :
124 : // 获取root
125 0 : u32 rootRank = 0;
126 0 : CHK_RET(GetRankByUserRank(commPlane, COMM_INDEX_0, param.root, rootRank));
127 :
128 0 : if (isUsedRegister) {
129 0 : PrepareData prepareData;
130 0 : SetPrepareData(prepareData, param, execMem, rootRank);
131 0 : CHK_RET(tempAlg->Prepare(prepareData));
132 0 : } else {
133 0 : CHK_RET(tempAlg->Prepare(execMem.inputMem, execMem.outputMem, execMem.outputMem, execMem.count,
134 : param.DataDes.dataType, param.stream, HCCL_REDUCE_RESERVED, rootRank));
135 : }
136 :
137 0 : CHK_RET(RunTemplate(tempAlg, combinedCommInfo));
138 :
139 0 : return HCCL_SUCCESS;
140 0 : }
141 :
142 : REGISTER_EXEC("BroadCastComm", BroadcastComm, CollBroadcastCommExecutor);
143 :
144 : } // namespace hccl
|