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