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_reduce_fast_double_ring_for_910_93_executor.h"
12 :
13 : namespace hccl {
14 :
15 4 : CollAllReduceFastDoubleRingFor91093Executor::CollAllReduceFastDoubleRingFor91093Executor(
16 4 : const HcclDispatcher dispatcher, std::unique_ptr<TopoMatcher>& topoMatcher)
17 4 : : CollAllReduceRingFor91093Executor(dispatcher, topoMatcher)
18 : {
19 4 : if (workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
20 4 : DMAReduceFlag_ = true;
21 : } else {
22 0 : DMAReduceFlag_ = false;
23 : }
24 4 : }
25 :
26 16 : HcclResult CollAllReduceFastDoubleRingFor91093Executor::DoubleRingReduceScatter(
27 : const std::string& tag, DeviceMem inputMem, DeviceMem outputMem, const u64 count, const HcclDataType dataType,
28 : const HcclReduceOp reductionOp, const std::vector<std::vector<Slice>> multRingsSliceZero, Stream stream,
29 : s32 profStage, const u64 baseOffset, const HcomCollOpInfo* opInfo,
30 : const std::vector<std::vector<Slice>> multRingsUserMemSlice, const bool disableDMAReduce)
31 : {
32 : (void)tag;
33 16 : HCCL_CONFIG_INFO(
34 : HCCL_ALG,
35 : "[CollAllReduceFastDoubleRingFor91093Executor][DoubleRingReduceScatter] DoubleRingReduceScatter starts");
36 16 : HcclResult ret = HCCL_SUCCESS;
37 16 : u32 ringNum = multRingsSliceZero.size();
38 16 : CHK_RET(CheckCommSize(COMM_LEVEL0, ringNum));
39 : // 拿到ring环映射关系
40 16 : SubCommInfo level0ZeroCommInfo = GetSubCommInfo(COMM_LEVEL0, COMM_INDEX_0);
41 16 : auto nicList = topoAttr_.nicList;
42 : std::vector<std::vector<u32>> multiRingsOrder
43 16 : = GetRingsOrderByTopoType(level0ZeroCommInfo.localRankSize, topoType_, nicList);
44 16 : u64 reduceAttr = GetReduceAttr(inputMem, outputMem, dataType, reductionOp);
45 16 : SubCommInfo level0RingCommInfo = GetSubCommInfo(COMM_LEVEL0, COMM_INDEX_0);
46 : // 生成两个ring上的userMemIn_上对应的slices
47 16 : std::vector<std::vector<Slice>> userMemInputSlicesOfDoubleRing;
48 16 : CHK_RET(CollectMultiRingsUserMemSlices(
49 : ringNum, dataType, opInfo, multRingsSliceZero, multiRingsOrder, multRingsUserMemSlice,
50 : userMemInputSlicesOfDoubleRing));
51 : // 生成两个ring上的rankOrder
52 16 : std::vector<std::vector<u32>> rankOrders;
53 16 : CHK_RET(CollectMultiRingsRankOrder(ringNum, multiRingsOrder, rankOrders));
54 : // 初始化executor
55 : std::unique_ptr<AlgTemplateBase> tempAlg
56 16 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_REDUCESCATTER_DB_RING_SLC, dispatcher_);
57 16 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_DB_RING_SLC in COMM_LEVEL0", __func__);
58 16 : CHK_SMART_PTR_NULL(tempAlg);
59 16 : ret = tempAlg->Prepare(
60 : inputMem, inputMem, outputMem, count, dataType, stream, multRingsSliceZero, reductionOp, LEVEL0_BRIDGE_RANK_ID,
61 16 : baseOffset, disableDMAReduce, reduceAttr, opInfo, topoAttr_.userRank, algResResp_->slaveStreams,
62 16 : algResResp_->notifiesMain, algResResp_->notifiesAux, rankOrders, userMemInputSlicesOfDoubleRing);
63 16 : CHK_PRT_RET(
64 : ret != HCCL_SUCCESS,
65 : HCCL_ERROR(
66 : "[CollAllReduceFastDoubleRingFor91093Executor][DoubleRingReduceScatter] Double ring "
67 : "ReduceScatter failed,return[%d]",
68 : ret),
69 : ret);
70 16 : u32 ringIndexOp = COMM_INDEX_0;
71 16 : u32 rankSize = level0RingCommInfo.localRankSize;
72 16 : ret = tempAlg->RegisterProfiler(
73 16 : ((ringIndexOp + 1) << PROF_RINGINDEX_OFFSET_OF_PLANEID) + (rankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID)
74 16 : + level0RingCommInfo.localRank,
75 : profStage, HCCL_EXEC_STEP_NOT_SET, stream);
76 16 : CHK_PRT_RET(
77 : ret != HCCL_SUCCESS,
78 : HCCL_ERROR(
79 : "[CollAllReduceFastDoubleRingFor91093Executor][DoubleRingReduceScatter] Double ring "
80 : "ReduceScatter failed,return[%d]",
81 : ret),
82 : ret);
83 : // 空拷贝用于后续操作附着
84 16 : CHK_RET(AlgTemplateBase::ExecEmptyTask(inputMem, outputMem, stream, dispatcher_));
85 16 : ret = RunTemplate(tempAlg, level0RingCommInfo);
86 16 : CHK_PRT_RET(
87 : ret != HCCL_SUCCESS,
88 : HCCL_ERROR(
89 : "[CollAllReduceFastDoubleRingFor91093Executor][DoubleRingReduceScatter] Double ring "
90 : "ReduceScatter failed,return[%d]",
91 : ret),
92 : ret);
93 :
94 16 : CHK_RET(AlgTemplateBase::ExecEmptyTask(inputMem, outputMem, stream, dispatcher_));
95 16 : return HCCL_SUCCESS;
96 16 : }
97 :
98 16 : HcclResult CollAllReduceFastDoubleRingFor91093Executor::DoubleRingAllGather(
99 : const std::string& tag, DeviceMem inputMem, DeviceMem outputMem, const u64 count, const HcclDataType dataType,
100 : const std::vector<std::vector<Slice>> multRingsSliceZero, Stream stream, s32 profStage, const u64 baseOffset,
101 : const HcomCollOpInfo* opInfo, const std::vector<std::vector<Slice>> multRingsUserMemSlice)
102 : {
103 : (void)tag;
104 16 : HCCL_CONFIG_INFO(
105 : HCCL_ALG, "[CollAllReduceFastDoubleRingFor91093Executor][DoubleRingAllGather] DoubleRingAllGather starts");
106 16 : HcclResult ret = HCCL_SUCCESS;
107 16 : u32 ringNum = multRingsSliceZero.size();
108 16 : CHK_RET(CheckCommSize(COMM_LEVEL0, ringNum));
109 : // 拿到ring环映射关系
110 16 : SubCommInfo level0ZeroCommInfo = GetSubCommInfo(COMM_LEVEL0, COMM_INDEX_0);
111 16 : auto nicList = topoAttr_.nicList;
112 : std::vector<std::vector<u32>> multiRingsOrder
113 16 : = GetRingsOrderByTopoType(level0ZeroCommInfo.localRankSize, topoType_, nicList);
114 : // 生成两个ring上的userMemOut_上对应的slices
115 16 : std::vector<std::vector<Slice>> userMemOutputSlicesOfDoubleRing;
116 16 : CHK_RET(CollectMultiRingsUserMemSlices(
117 : ringNum, dataType, opInfo, multRingsSliceZero, multiRingsOrder, multRingsUserMemSlice,
118 : userMemOutputSlicesOfDoubleRing));
119 : // 生成两个ring上的rankOrder
120 16 : std::vector<std::vector<u32>> rankOrders;
121 16 : CHK_RET(CollectMultiRingsRankOrder(ringNum, multiRingsOrder, rankOrders));
122 : // 初始化executor
123 16 : std::unique_ptr<AlgTemplateBase> tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
124 16 : TemplateType::TEMPLATE_ALIGNED_ALL_GATHER_DOUBLE_RING, dispatcher_);
125 16 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALIGNED_ALL_GATHER_DOUBLE_RING in COMM_LEVEL0", __func__);
126 16 : CHK_SMART_PTR_NULL(tempAlg);
127 16 : CHK_RET(tempAlg->Prepare(
128 : const_cast<HcomCollOpInfo*>(opInfo), topoAttr_.userRank, algResResp_->slaveStreams, algResResp_->notifiesMain,
129 : algResResp_->notifiesAux, rankOrders, userMemOutputSlicesOfDoubleRing));
130 :
131 16 : ret = tempAlg->Prepare(
132 : outputMem, outputMem, inputMem, count, dataType, stream, multRingsSliceZero, HCCL_REDUCE_RESERVED,
133 : LEVEL0_BRIDGE_RANK_ID, baseOffset);
134 16 : CHK_PRT_RET(
135 : ret != HCCL_SUCCESS,
136 : HCCL_ERROR(
137 : "[CollAllReduceFastDoubleRingFor91093Executor][DoubleRingAllGather]Double ring "
138 : "AllGather failed, return[%d]",
139 : ret),
140 : ret);
141 16 : u32 ringIndexOp = COMM_INDEX_0;
142 16 : u32 rankSize = level0ZeroCommInfo.localRankSize;
143 16 : ret = tempAlg->RegisterProfiler(
144 16 : ((ringIndexOp + 1) << PROF_RINGINDEX_OFFSET_OF_PLANEID) + (rankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID)
145 16 : + level0ZeroCommInfo.localRank,
146 : profStage, HCCL_EXEC_STEP_NOT_SET, stream);
147 16 : CHK_PRT_RET(
148 : ret != HCCL_SUCCESS,
149 : HCCL_ERROR(
150 : "[CollAllReduceFastDoubleRingFor91093Executor][DoubleRingAllGather]Double ring "
151 : "AllGather failed, return[%d]",
152 : ret),
153 : ret);
154 :
155 : // 空拷贝用于后续操作附着
156 16 : CHK_RET(AlgTemplateBase::ExecEmptyTask(inputMem, outputMem, stream, dispatcher_));
157 16 : ret = RunTemplate(tempAlg, level0ZeroCommInfo);
158 16 : CHK_PRT_RET(
159 : ret != HCCL_SUCCESS,
160 : HCCL_ERROR(
161 : "[CollAllReduceFastDoubleRingFor91093Executor][DoubleRingAllGather] Double ring "
162 : "AllGather failed,return[%d]",
163 : ret),
164 : ret);
165 : // 添加空task,保证执行时不乱序
166 16 : CHK_RET(AlgTemplateBase::ExecEmptyTask(inputMem, outputMem, stream, dispatcher_));
167 16 : return HCCL_SUCCESS;
168 16 : }
169 :
170 16 : HcclResult CollAllReduceFastDoubleRingFor91093Executor::RunIntraSeverReduceScatter(
171 : const std::string& tag, DeviceMem& inputMem, DeviceMem& outputMem, const u64 count, const HcclDataType& dataType,
172 : const HcclReduceOp& reductionOp, const std::vector<std::vector<Slice>>& multRingsSliceZero, const Stream& stream,
173 : s32 profStage, const u64 baseOffset, const HcomCollOpInfo* opInfo,
174 : const std::vector<std::vector<Slice>>& multRingsUserMemSlice, const bool disableDMAReduce)
175 : {
176 16 : CHK_RET(DoubleRingReduceScatter(
177 : tag, inputMem, outputMem, count, dataType, reductionOp, multRingsSliceZero, stream, profStage, baseOffset,
178 : opInfo, multRingsUserMemSlice, disableDMAReduce));
179 16 : return HCCL_SUCCESS;
180 : }
181 :
182 16 : HcclResult CollAllReduceFastDoubleRingFor91093Executor::RunIntraSeverAllGather(
183 : const std::string& tag, DeviceMem& inputMem, DeviceMem& outputMem, const u64 count, const HcclDataType& dataType,
184 : const std::vector<std::vector<Slice>>& multRingsSliceZero, const Stream& stream, s32 profStage,
185 : const u64 baseOffset, const HcomCollOpInfo* opInfo, const std::vector<std::vector<Slice>>& multRingsUserMemSlice)
186 : {
187 16 : CHK_RET(DoubleRingAllGather(
188 : tag, inputMem, outputMem, count, dataType, multRingsSliceZero, stream, profStage, baseOffset, opInfo,
189 : multRingsUserMemSlice));
190 16 : return HCCL_SUCCESS;
191 : }
192 :
193 : REGISTER_EXEC(
194 : "AllReduceFastDoubleRingFor91093Executor", AllReduceFastDoubleRingFor91093,
195 : CollAllReduceFastDoubleRingFor91093Executor);
196 :
197 : } // namespace hccl
|