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