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 "send_receive.h"
12 :
13 : namespace hccl {
14 0 : SendReceive::SendReceive(
15 : const HcclDispatcher dispatcher,
16 : const std::shared_ptr<Transport> &link,
17 : const u32 peerRank,
18 : const u64 chunkNum,
19 0 : bool retryEnable)
20 : : AlgTemplateBase(dispatcher),
21 0 : transLink_(link),
22 0 : peerRank_(peerRank),
23 0 : chunkSize_(chunkNum),
24 0 : retryEnable_(retryEnable)
25 : {
26 0 : }
27 :
28 0 : SendReceive::~SendReceive()
29 : {
30 0 : }
31 :
32 0 : HcclResult SendReceive::SendPrepare(
33 : const DeviceMem &inputMem,
34 : const u32 destRank,
35 : const Stream &stream)
36 : {
37 : /* 参数赋值 */
38 0 : inputMem_ = inputMem;
39 0 : stream_ = stream;
40 0 : peerRank_ = destRank;
41 0 : return HCCL_SUCCESS;
42 : }
43 :
44 0 : HcclResult SendReceive::ReceivePrepare(
45 : const DeviceMem &outputMem,
46 : const u32 srcRank,
47 : const Stream &stream)
48 : {
49 : /* 参数赋值 */
50 0 : outputMem_ = outputMem;
51 0 : stream_ = stream;
52 0 : peerRank_ = srcRank;
53 :
54 0 : return HCCL_SUCCESS;
55 : }
56 :
57 0 : HcclResult SendReceive::SendRunAsync()
58 : {
59 0 : if (!inputMem_) {
60 0 : HCCL_ERROR("[SendReceive][SendRunAsync]SendRunAsync inputmem is null");
61 0 : return HCCL_E_PTR;
62 : }
63 :
64 0 : CHK_SMART_PTR_NULL(transLink_);
65 0 : u64 sizePerRound = 0;
66 0 : u64 sizePerSlice = chunkSize_;
67 0 : u64 length = inputMem_.size();
68 0 : u64 offset = 0;
69 0 : for (u64 sizeResidue = length; sizeResidue > 0; sizeResidue -= sizePerRound) {
70 0 : HcclResult ret = transLink_->TxAck(stream_);
71 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[SendReceive][SendRunAsync]tx ack run failed"), ret);
72 0 : ret = transLink_->RxAck(stream_);
73 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[SendReceive][SendRunAsync]rx ack run failed"), ret);
74 :
75 0 : offset += sizePerRound;
76 0 : sizePerRound = (sizeResidue > sizePerSlice) ? sizePerSlice : sizeResidue;
77 0 : void* localAddr = static_cast<u8 *>(inputMem_.ptr()) + offset;
78 0 : HCCL_DEBUG("tx async inputmem's offset[%llu] size[%llu]", offset, sizePerRound);
79 :
80 0 : ret = transLink_->TxAsync(UserMemType::OUTPUT_MEM, offset, localAddr, sizePerRound, stream_);
81 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[SendReceive][SendRunAsync]tx async offset[%llu] "\
82 : "size[%llu] failed", offset, sizePerRound), ret);
83 :
84 0 : ret = transLink_->RxAsync(UserMemType::OUTPUT_MEM, 0, nullptr, 0, stream_);
85 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[SendReceive][ReceiveRunAsync]tx async failed"), ret);
86 :
87 0 : ret = transLink_->RxWaitDone(stream_);
88 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[SendReceive][SendRunAsync]RxWaitDone failed"), ret);
89 0 : ret = transLink_->TxWaitDone(stream_);
90 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[SendReceive][SendRunAsync]TxWaitDone failed"), ret);
91 : }
92 0 : return HCCL_SUCCESS;
93 : }
94 :
95 0 : HcclResult SendReceive::ReceiveRunAsync()
96 : {
97 0 : if (!outputMem_) {
98 0 : HCCL_ERROR("[SendReceive][ReceiveRunAsync]ReceiveRunAsync outputmem is null");
99 0 : return HCCL_E_PTR;
100 : }
101 0 : CHK_SMART_PTR_NULL(transLink_);
102 :
103 0 : HCCL_DEBUG("[SendReceive][ReceiveRunAsync]ReceiveRunAsync begins");
104 0 : u64 sizePerRound = 0;
105 0 : u64 sizePerSlice = chunkSize_;
106 0 : u64 length = outputMem_.size();
107 0 : u64 offset = 0;
108 0 : for (u64 sizeResidue = length; sizeResidue > 0; sizeResidue -= sizePerRound) {
109 0 : HcclResult ret = transLink_->TxAck(stream_);
110 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[SendReceive][ReceiveRunAsync]tx ack failed"), ret);
111 0 : ret = transLink_->RxAck(stream_);
112 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[SendReceive][ReceiveRunAsync]rx ack failed"), ret);
113 :
114 0 : offset += sizePerRound;
115 0 : sizePerRound = (sizeResidue > sizePerSlice) ? sizePerSlice : sizeResidue;
116 0 : void* localAddr = static_cast<u8 *>(outputMem_.ptr()) + offset;
117 0 : HCCL_DEBUG("rx async outputmem's offset[%llu] size[%llu]", offset, sizePerRound);
118 :
119 0 : ret = transLink_->RxAsync(UserMemType::OUTPUT_MEM, offset, localAddr, sizePerRound, stream_);
120 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[SendReceive][ReceiveRunAsync]rx async with offset[%llu] "\
121 : "size[%llu] failed", offset, sizePerRound), ret);
122 :
123 0 : ret = transLink_->TxAsync(UserMemType::OUTPUT_MEM, 0, nullptr, 0, stream_);
124 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[SendReceive][ReceiveRunAsync]tx async failed"), ret);
125 : }
126 0 : return HCCL_SUCCESS;
127 : }
128 :
129 0 : HcclResult SendReceive::BatchSendRunAsync()
130 : {
131 0 : if (!inputMem_) {
132 0 : HCCL_ERROR("[SendReceive][BatchSendRunAsync] BatchSendRunAsync inputmem is null");
133 0 : return HCCL_E_PTR;
134 : }
135 :
136 0 : CHK_SMART_PTR_NULL(transLink_);
137 0 : u64 sizePerRound = 0;
138 0 : u64 sizePerSlice = chunkSize_;
139 0 : u64 length = inputMem_.size();
140 0 : u64 offset = 0;
141 0 : for (u64 sizeResidue = length; sizeResidue > 0; sizeResidue -= sizePerRound) {
142 0 : HcclResult ret = transLink_->TxPrepare(stream_);
143 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[SendReceive][BatchSendRunAsync]tx ack run failed"), ret);
144 :
145 0 : offset += sizePerRound;
146 0 : sizePerRound = (sizeResidue > sizePerSlice) ? sizePerSlice : sizeResidue;
147 0 : void* localAddr = static_cast<u8 *>(inputMem_.ptr()) + offset;
148 0 : HCCL_DEBUG("tx async inputmem's offset[%llu] size[%llu]", offset, sizePerRound);
149 :
150 0 : ret = transLink_->TxData(UserMemType::OUTPUT_MEM, offset, localAddr, sizePerRound, stream_);
151 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[SendReceive][BatchSendRunAsync]tx async offset[%llu] "\
152 : "size[%llu] failed", offset, sizePerRound), ret);
153 :
154 0 : ret = transLink_->TxDone(stream_);
155 :
156 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[SendReceive][BatchSendRunAsync]TxWaitDone failed"), ret);
157 : }
158 0 : return HCCL_SUCCESS;
159 : }
160 :
161 0 : HcclResult SendReceive::BatchReceiveRunAsync()
162 : {
163 0 : if (!outputMem_) {
164 0 : HCCL_ERROR("[SendReceive][ReceiveRunAsync]ReceiveRunAsync outputmem is null");
165 0 : return HCCL_E_PTR;
166 : }
167 0 : CHK_SMART_PTR_NULL(transLink_);
168 :
169 0 : u64 sizePerRound = 0;
170 0 : u64 sizePerSlice = chunkSize_;
171 0 : u64 length = outputMem_.size();
172 0 : u64 offset = 0;
173 0 : for (u64 sizeResidue = length; sizeResidue > 0; sizeResidue -= sizePerRound) {
174 0 : HcclResult ret = transLink_->RxPrepare(stream_);
175 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[SendReceive][BatchReceiveRunAsync]rx ack failed"), ret);
176 :
177 0 : offset += sizePerRound;
178 0 : sizePerRound = (sizeResidue > sizePerSlice) ? sizePerSlice : sizeResidue;
179 0 : void* localAddr = static_cast<u8 *>(outputMem_.ptr()) + offset;
180 0 : HCCL_DEBUG("rx async outputmem's offset[%llu] size[%llu]", offset, sizePerRound);
181 :
182 0 : ret = transLink_->RxData(UserMemType::INPUT_MEM, offset, localAddr, sizePerRound, stream_);
183 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[SendReceive][BatchReceiveRunAsync]rx async with offset[%llu] "\
184 : "size[%llu] failed", offset, sizePerRound), ret);
185 :
186 0 : ret = transLink_->RxDone(stream_);
187 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[SendReceive][BatchReceiveRunAsync]TxDataSignal offset[%llu]"\
188 : "size[%llu] failed", offset, sizePerRound), ret);
189 : }
190 0 : return HCCL_SUCCESS;
191 : }
192 : } // namespace hccl
193 :
|