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