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_reduce_scatter_fast_double_ring_for_910_93_executor.h"
12 :
13 : namespace hccl {
14 2 : CollReduceScatterFastDoubleRingFor91093Executor::CollReduceScatterFastDoubleRingFor91093Executor(const HcclDispatcher dispatcher,
15 2 : std::unique_ptr<TopoMatcher> &topoMatcher)
16 2 : : CollAlignedReduceScatterDoubleRingFor91093Executor(dispatcher, topoMatcher)
17 : {
18 2 : DMAReduceFlag_ = (workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE);
19 2 : }
20 :
21 17 : HcclResult CollReduceScatterFastDoubleRingFor91093Executor::DoubleRingReduceScatter(const std::string &tag, DeviceMem inputMem, DeviceMem outputMem,
22 : const u64 count, const HcclDataType dataType, const HcclReduceOp reductionOp,
23 : const std::vector<std::vector<Slice> > multRingsSliceZero, Stream stream, s32 profStage,
24 : const u64 baseOffset, const HcomCollOpInfo *opInfo,
25 : const std::vector<std::vector<Slice>> multRingsUserMemSlice, const bool disableDMAReduce)
26 : {
27 : (void)tag;
28 17 : HCCL_CONFIG_INFO(HCCL_ALG,
29 : "[CollReduceScatterFastDoubleRingFor91093Executor][DoubleRingReduceScatter] DoubleRingReduceScatter starts");
30 17 : HcclResult ret = HCCL_SUCCESS;
31 17 : u32 ringNum = multRingsSliceZero.size();
32 17 : CHK_RET(CheckCommSize(COMM_LEVEL0, ringNum));
33 :
34 : // 拿到ring环映射关系
35 17 : SubCommInfo level0ZeroCommInfo = GetSubCommInfo(COMM_LEVEL0, COMM_INDEX_0);
36 17 : auto nicList = topoAttr_.nicList;
37 : std::vector<std::vector<u32>> multiRingsOrder =
38 17 : GetRingsOrderByTopoType(level0ZeroCommInfo.localRankSize, topoType_, nicList);
39 :
40 17 : u64 reduceAttr = GetReduceAttr(inputMem, outputMem, dataType, reductionOp);
41 :
42 17 : SubCommInfo level0RingCommInfo = GetSubCommInfo(COMM_LEVEL0, COMM_INDEX_0);
43 : // 生成两个ring上的userMemIn_上对应的slices
44 17 : std::vector<std::vector<Slice>> userMemInputSlicesOfDoubleRing;
45 17 : CHK_RET(CollectMultiRingsUserMemSlices(ringNum, dataType,
46 : opInfo, multRingsSliceZero,
47 : multiRingsOrder, multRingsUserMemSlice,
48 : userMemInputSlicesOfDoubleRing));
49 : // 生成两个ring上的rankOrder
50 16 : std::vector<std::vector<u32>> rankOrders;
51 16 : CHK_RET(CollectMultiRingsRankOrder(ringNum, multiRingsOrder, rankOrders));
52 : // 初始化executor
53 16 : std::unique_ptr<AlgTemplateBase> tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
54 16 : TemplateType::TEMPLATE_REDUCESCATTER_DB_RING_SLC, dispatcher_);
55 16 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_DB_RING_SLC in COMM_LEVEL0", __func__);
56 16 : CHK_SMART_PTR_NULL(tempAlg);
57 16 : ret = tempAlg->Prepare(inputMem, inputMem, outputMem, count, dataType, stream, multRingsSliceZero,
58 : reductionOp, LEVEL0_BRIDGE_RANK_ID, baseOffset, disableDMAReduce, reduceAttr, opInfo,
59 16 : topoAttr_.userRank, algResResp_->slaveStreams, algResResp_->notifiesMain, algResResp_->notifiesAux,
60 : rankOrders, userMemInputSlicesOfDoubleRing);
61 16 : CHK_PRT_RET(ret != HCCL_SUCCESS,
62 : HCCL_ERROR("[CollReduceScatterFastDoubleRingFor91093Executor][DoubleRingReduceScatter] Double ring ReduceScatter failed"
63 : "failed,return[%d]", ret), ret);
64 16 : u32 ringIndexOp = COMM_INDEX_0;
65 16 : u32 rankSize = level0RingCommInfo.localRankSize;
66 16 : ret = tempAlg->RegisterProfiler(
67 16 : ((ringIndexOp + 1) << PROF_RINGINDEX_OFFSET_OF_PLANEID) +
68 16 : (rankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID) + level0RingCommInfo.localRank,
69 : profStage, HCCL_EXEC_STEP_NOT_SET, stream);
70 16 : CHK_PRT_RET(ret != HCCL_SUCCESS,
71 : HCCL_ERROR("[CollReduceScatterFastDoubleRingFor91093Executor][DoubleRingReduceScatter] Double ring ReduceScatter failed "
72 : "failed,return[%d]", ret), ret);
73 : // 空拷贝用于后续操作附着
74 16 : CHK_RET(AlgTemplateBase::ExecEmptyTask(inputMem, outputMem, stream, dispatcher_));
75 16 : ret = RunTemplate(tempAlg, level0RingCommInfo);
76 16 : CHK_PRT_RET(ret != HCCL_SUCCESS,
77 : HCCL_ERROR("[CollReduceScatterFastDoubleRingFor91093Executor][DoubleRingReduceScatter] Double ring ReduceScatter failed "
78 : "failed,return[%d]", ret), ret);
79 :
80 16 : CHK_RET(AlgTemplateBase::ExecEmptyTask(inputMem, outputMem, stream, dispatcher_));
81 16 : return HCCL_SUCCESS;
82 17 : }
83 : REGISTER_EXEC("ReduceScatterFastDoubleRingFor91093Executor", ReduceScatterFastDoubleRingFor91093, CollReduceScatterFastDoubleRingFor91093Executor);
84 : }
|