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