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