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 "comm_ring.h"
12 :
13 : namespace hccl {
14 10 : CommRing::CommRing(const std::string &collectiveId, const u32 userRank,
15 : const u32 userRankSize, const u32 rank, const u32 rankSize, const TopoType topoFlag,
16 : const HcclDispatcher dispatcher, const std::unique_ptr<NotifyPool> ¬ifyPool,
17 : std::map<HcclIpAddress, HcclNetDevCtx> &netDevCtxMap,
18 : const IntraExchanger &exchanger, const std::vector<RankInfo> paraVector,
19 : const DeviceMem& inputMem, const DeviceMem& outputMem, const bool isUsedRdmaLevel0,
20 : const std::string &tag,
21 : const NICDeployment nicDeployInner, const bool useOneDoorbell,
22 10 : const bool isAicpuModeEn, const bool isHaveCpuRank, const bool useSuperPodMode)
23 : : CommBase(collectiveId, userRank, userRankSize, rank, rankSize, paraVector, topoFlag, dispatcher, notifyPool,
24 : netDevCtxMap, exchanger, inputMem, outputMem, isUsedRdmaLevel0, tag, nicDeployInner, 0, useOneDoorbell, isAicpuModeEn, INVALID_UINT,
25 10 : isHaveCpuRank, useSuperPodMode)
26 : {
27 10 : }
28 :
29 19 : CommRing::~CommRing()
30 : {
31 19 : }
32 :
33 0 : HcclResult CommRing::CalcLink()
34 : {
35 0 : u32 dstClientRank = INVALID_VALUE_RANKID;
36 0 : u32 dstServerRank = INVALID_VALUE_RANKID;
37 0 : HcclResult ret = HCCL_SUCCESS;
38 0 : if (rank_ == HCCL_RANK_ZERO) { // 当前rank为rank0
39 : // rank 作为server
40 0 : dstClientRank = rank_ + HCCL_RANK_OFFSET;
41 0 : ret = CalcLinksNum(MachineType::MACHINE_SERVER_TYPE, dstClientRank);
42 0 : CHK_PRT_RET(ret != HCCL_SUCCESS,
43 : HCCL_ERROR("[Calc][Link]comm ring calc links num failed, type[%d], dstClientRank[%u]",
44 : static_cast<int32_t>(MachineType::MACHINE_SERVER_TYPE), dstClientRank), ret);
45 :
46 0 : if (rankSize_ > HCCL_RANK_SIZE_EQ_TWO) {
47 : // rank 作为client
48 0 : dstServerRank = rankSize_ - HCCL_RANK_OFFSET;
49 0 : ret = CalcLinksNum(MachineType::MACHINE_CLIENT_TYPE, dstServerRank);
50 0 : CHK_PRT_RET(ret != HCCL_SUCCESS,
51 : HCCL_ERROR("[Calc][Link]comm ring calc links num failed, type[%d], dstServerRank[%u]",
52 : static_cast<int32_t>(MachineType::MACHINE_CLIENT_TYPE), dstServerRank), ret);
53 : }
54 0 : } else if ((rankSize_ - HCCL_RANK_OFFSET) == rank_) { // 当前rank为ring环尾,rankx(x = (rankSize_ - 1))
55 0 : if (rankSize_ > HCCL_RANK_SIZE_EQ_TWO) {
56 : // rank 作为server
57 0 : dstClientRank = HCCL_RANK_ZERO;
58 0 : ret = CalcLinksNum(MachineType::MACHINE_SERVER_TYPE, dstClientRank);
59 0 : CHK_PRT_RET(ret != HCCL_SUCCESS,
60 : HCCL_ERROR("[Calc][Link]comm ring calc links num failed, type[%d], dstClientRank[%u]",
61 : static_cast<int32_t>(MachineType::MACHINE_SERVER_TYPE), dstClientRank), ret);
62 : }
63 :
64 : // rank 作为client
65 0 : dstServerRank = rank_ - HCCL_RANK_OFFSET;
66 0 : ret = CalcLinksNum(MachineType::MACHINE_CLIENT_TYPE, dstServerRank);
67 0 : CHK_PRT_RET(ret != HCCL_SUCCESS,
68 : HCCL_ERROR("[Calc][Link]comm ring calc links num failed, type[%d], dstServerRank[%u]",
69 : static_cast<int32_t>(MachineType::MACHINE_CLIENT_TYPE), dstServerRank), ret);
70 : } else { // 奇数先创建client,偶数先创建server
71 0 : if ((rank_ % 2) != 0) { // 模2判断奇偶性,rank为奇数
72 : // rank 作为client
73 0 : dstServerRank = rank_ - HCCL_RANK_OFFSET;
74 0 : ret = CalcLinksNum(MachineType::MACHINE_CLIENT_TYPE, dstServerRank);
75 0 : CHK_PRT_RET(ret != HCCL_SUCCESS,
76 : HCCL_ERROR("[Calc][Link]comm ring calc links num failed, type[%d], dstServerRank[%u]",
77 : static_cast<int32_t>(MachineType::MACHINE_CLIENT_TYPE), dstServerRank), ret);
78 :
79 : // rank 作为server
80 0 : dstClientRank = rank_ + HCCL_RANK_OFFSET;
81 0 : ret = CalcLinksNum(MachineType::MACHINE_SERVER_TYPE, dstClientRank);
82 0 : CHK_PRT_RET(ret != HCCL_SUCCESS,
83 : HCCL_ERROR("[Calc][Link]comm ring calc links num failed, type[%d], dstClientRank[%u]",
84 : static_cast<int32_t>(MachineType::MACHINE_SERVER_TYPE), dstClientRank), ret);
85 : } else { // rank为偶数
86 : // rank 作为server
87 0 : dstClientRank = rank_ + HCCL_RANK_OFFSET;
88 0 : ret = CalcLinksNum(MachineType::MACHINE_SERVER_TYPE, dstClientRank);
89 0 : CHK_PRT_RET(ret != HCCL_SUCCESS,
90 : HCCL_ERROR("[Calc][Link]comm ring calc links num failed, type[%d], dstClientRank[%u]",
91 : static_cast<int32_t>(MachineType::MACHINE_SERVER_TYPE), dstClientRank), ret);
92 :
93 : // rank 作为client
94 0 : dstServerRank = rank_ - HCCL_RANK_OFFSET;
95 0 : ret = CalcLinksNum(MachineType::MACHINE_CLIENT_TYPE, dstServerRank);
96 0 : CHK_PRT_RET(ret != HCCL_SUCCESS,
97 : HCCL_ERROR("[Calc][Link]comm ring calc links num failed, type[%d], dstServerRank[%u]",
98 : static_cast<int32_t>(MachineType::MACHINE_CLIENT_TYPE), dstServerRank), ret);
99 : }
100 : }
101 :
102 0 : return HCCL_SUCCESS;
103 : }
104 :
105 : // 获取每个 link 需要的 socket 数量
106 2 : u32 CommRing::GetSocketsPerLink()
107 : {
108 2 : const u32 rdmaTaskNumRatio = 4; // server间ring算法每个link上rdma task数为 4*rank size
109 2 : HcclWorkflowMode workFlowMode = GetWorkflowMode();
110 2 : if (workFlowMode != HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
111 0 : return (rankSize_ * rdmaTaskNumRatio + (HCCP_SQ_TEMPLATE_CAPACITY - 1)) / HCCP_SQ_TEMPLATE_CAPACITY;
112 : } else {
113 : // op base 场景每个link使用 1 个QP,只需要建立1个socket链接
114 2 : return 1;
115 : }
116 : }
117 :
118 0 : void CommRing::SetMachineLinkMode(MachinePara &machinePara)
119 : {
120 0 : machinePara.linkMode = LinkMode::LINK_DUPLEX_MODE;
121 0 : }
122 : } // namespace hccl
123 :
|