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 <math.h>
12 : #include "alg_template_register.h"
13 : #include "all_reduce_doubling_local_reduce.h"
14 :
15 : namespace hccl {
16 :
17 : // Doubling算法实现AllReduce,只用于server内通信
18 0 : AllReduceDoublingLocalReduce::AllReduceDoublingLocalReduce(const HcclDispatcher dispatcher)
19 0 : : AlgTemplateBase(dispatcher)
20 0 : {}
21 :
22 0 : AllReduceDoublingLocalReduce::~AllReduceDoublingLocalReduce() {}
23 :
24 0 : HcclResult AllReduceDoublingLocalReduce::Prepare(u64 reduceAttrBitMap, [[maybe_unused]] HcomCollOpInfo* opInfo)
25 : {
26 0 : reduceAttr_ = reduceAttrBitMap;
27 0 : return HCCL_SUCCESS;
28 : }
29 :
30 0 : HcclResult AllReduceDoublingLocalReduce::RunAsync(const u32 rank, const u32 rankSize, const std::vector<LINK>& links)
31 : {
32 0 : HCCL_INFO(
33 : "[AllReduceDoublingLocalReduce] runAsync rank[%u] ranksize[%u] inputMem[%p] outputMem[%p] count[%llu]", rank,
34 : rankSize, inputMem_.ptr(), outputMem_.ptr(), count_);
35 :
36 : // 基本的检查
37 0 : CHK_RET(SimpleCheck(rank, rankSize, links));
38 :
39 : // 判断rank_size == 1
40 0 : if (rankSize == 1) {
41 : // 对于Doubling,input和output必须是两块不同的内存
42 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, outputMem_, inputMem_, stream_));
43 0 : return HCCL_SUCCESS;
44 : }
45 :
46 : // 设置Slices
47 0 : if (slices_.size() != 0) {
48 0 : HCCL_WARNING("[AllReduceDoublingLocalReduce] slices_ will be not used in executor.");
49 : }
50 :
51 : // 执行算法
52 0 : CHK_RET(RunAllReduce(rank, rankSize, links));
53 :
54 0 : HCCL_INFO("AllReduceDoublingLocalReduce finished: rank[%u] ranksize[%u]", rank, rankSize);
55 0 : return HCCL_SUCCESS;
56 : }
57 :
58 0 : HcclResult AllReduceDoublingLocalReduce::SimpleCheck(const u32 rank, const u32 rankSize, const std::vector<LINK>& links)
59 : {
60 : // 判断stream, dispatcher是否为空
61 0 : CHK_SMART_PTR_NULL(dispatcher_);
62 0 : CHK_PTR_NULL(stream_.ptr());
63 :
64 : // 判断Memory是否为空
65 0 : CHK_PRT_RET(!inputMem_, HCCL_ERROR("[AllReduceDoublingLocalReduce] rank[%u] inputmem is null", rank), HCCL_E_PTR);
66 0 : CHK_PRT_RET(!outputMem_, HCCL_ERROR("[AllReduceDoublingLocalReduce] rank[%u] outputmem is null", rank), HCCL_E_PTR);
67 :
68 : // 必须有两块memory
69 0 : CHK_PRT_RET(
70 : inputMem_ == outputMem_,
71 : HCCL_ERROR(
72 : "[AllReduceDoublingLocalReduce] rank[%u] inputMem and outputMem "
73 : "should be different",
74 : rank),
75 : HCCL_E_PARA);
76 :
77 : // 判断links数量是否正确
78 0 : CHK_PRT_RET(
79 : links.size() < rankSize,
80 : HCCL_ERROR(
81 : "[AllReduceDoublingLocalReduce] rank[%u] link size[%llu] is "
82 : "less than rank size[%u]",
83 : rank, links.size(), rankSize),
84 : HCCL_E_PARA);
85 :
86 : // 判断rankSize是否为2的幂次
87 0 : CHK_PRT_RET(
88 : (rankSize & (rankSize - 1)) != 0,
89 : HCCL_ERROR(
90 : "[AllReduceDoublingLocalReduce] "
91 : "rankSize must be power of 2, but get rankSize=%u",
92 : rankSize),
93 : HCCL_E_PARA);
94 0 : return HCCL_SUCCESS;
95 : }
96 :
97 : HcclResult
98 0 : AllReduceDoublingLocalReduce::RunAllReduce(const u32 rank, const u32 rankSize, const std::vector<LINK>& links)
99 : {
100 0 : u64 totalSize = count_ * SIZE_TABLE[dataType_];
101 0 : DeviceMem localCclInMem = inputMem_.range(0, totalSize);
102 0 : DeviceMem localCclOutMem = outputMem_.range(0, totalSize);
103 :
104 0 : u32 nSteps = static_cast<u32>(log2(rankSize));
105 0 : for (u32 step = 0; step < nSteps; step++) {
106 : // 计算邻居并获取link
107 0 : u32 neighbor = rank ^ (1 << step);
108 0 : const LINK& link = links[neighbor];
109 0 : CHK_PTR_NULL(link);
110 0 : if (link->GetLinkType() == LinkType::LINK_ROCE) {
111 0 : CHK_RET(RunAllReduceRDMA(link, localCclInMem, localCclOutMem));
112 : } else {
113 0 : CHK_RET(RunAllReduceSDMA(link, localCclOutMem, totalSize));
114 : }
115 :
116 : // 从本端的cclout Reduce到本端的cclIn
117 0 : CHK_RET(HcclReduceAsync(
118 : dispatcher_, localCclOutMem.ptr(), count_, dataType_, reductionOp_, stream_, localCclInMem.ptr(),
119 : INVALID_VALUE_RANKID, LinkType::LINK_ONCHIP, reduceAttr_));
120 : }
121 0 : return HCCL_SUCCESS;
122 0 : }
123 :
124 0 : HcclResult AllReduceDoublingLocalReduce::RunAllReduceSDMA(const LINK& link, DeviceMem& localCclOutMem, u64 totalSize)
125 : {
126 : // 前同步
127 0 : CHK_RET(link->TxAck(stream_));
128 0 : CHK_RET(link->RxAck(stream_));
129 :
130 : // 从对端的cclIn读到本端的cclOut
131 0 : void* remMemPtr = nullptr;
132 0 : CHK_RET(link->GetRemoteMem(UserMemType::INPUT_MEM, &remMemPtr));
133 0 : DeviceMem remoteCclInMem = DeviceMem::create(remMemPtr, totalSize);
134 0 : CHK_RET(HcclD2DMemcpyAsync(
135 : dispatcher_, localCclOutMem, remoteCclInMem, stream_, link->GetRemoteRank(), link->GetLinkType()));
136 :
137 : // 尾同步
138 0 : CHK_RET(link->TxDataSignal(stream_));
139 0 : CHK_RET(link->RxDataSignal(stream_));
140 0 : return HCCL_SUCCESS;
141 0 : }
142 :
143 : HcclResult
144 0 : AllReduceDoublingLocalReduce::RunAllReduceRDMA(const LINK& link, DeviceMem& localCclInMem, DeviceMem& localCclOutMem)
145 : {
146 : // 前同步
147 0 : CHK_RET(link->TxAck(stream_));
148 0 : CHK_RET(link->RxAck(stream_));
149 :
150 : // RdmaSend + Record
151 0 : CHK_RET(link->TxAsync(UserMemType::OUTPUT_MEM, 0, localCclInMem.ptr(), localCclInMem.size(), stream_));
152 : // wait
153 0 : CHK_RET(link->RxAsync(UserMemType::INPUT_MEM, 0, localCclOutMem.ptr(), localCclOutMem.size(), stream_));
154 : // 后同步
155 0 : CHK_RET(link->PostFinAck(stream_));
156 0 : CHK_RET(link->WaitFinAck(stream_));
157 0 : return HCCL_SUCCESS;
158 : }
159 : REGISTER_TEMPLATE(TemplateType::TEMPLATE_ALL_REDUCE_DOUBLING_LOCAL_REDUCE, AllReduceDoublingLocalReduce);
160 : } // namespace hccl
|