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 <ios>
12 : #include <iostream>
13 :
14 : #include "log.h"
15 :
16 : #include "ccu_temp_all_to_all_mesh_1D.h"
17 : #include "alg_data_trans_wrapper.h"
18 : #include "ccu_instruction_all_to_all_mesh1d.h"
19 : #include "ccu_rank_group.h"
20 : #include "ccu_ctx_creator_registry.h"
21 : #include "ccu_context_all_to_all_mesh1d.h"
22 :
23 : namespace Hccl {
24 :
25 : static CcuInstRegister<CcuContextAllToAllMesh1D> registrarAllToAll(CcuInstType::CCU_ALLTOALL_MESH_1D_DIRECT);
26 :
27 0 : CcuTempAllToAllMesh1D::CcuTempAllToAllMesh1D(const RankId virtualRank, const u32 tempRankSize,
28 : const std::vector<std::vector<RankId>> &tempVTopo,
29 0 : const std::map<RankId, u32> &tempVirtRankMap)
30 0 : : CcuAlgTemplateBase(virtualRank, tempRankSize, tempVTopo, tempVirtRankMap)
31 : {
32 0 : }
33 :
34 0 : CcuTempAllToAllMesh1D::~CcuTempAllToAllMesh1D()
35 : {
36 0 : }
37 :
38 0 : HcclResult CcuTempAllToAllMesh1D::CalcRes(AlgTempResReq &tempResReq)
39 : {
40 0 : tempResReq.queNum = 1;
41 0 : tempResReq.streamNum = tempResReq.queNum;
42 0 : HCCL_INFO("[CalcRes] tempResReq.queNum[%u]", tempResReq.queNum);
43 0 : CHK_RET(CalcResLinksMesh(myRank_, tempRankSize_, tempVTopo_, linkNumBtwPeers_, tempResReq));
44 0 : return HcclResult::HCCL_SUCCESS;
45 : }
46 : /*
47 : dataSize / (rankSize) --> chunkSize
48 : dataSize / (rankSize * queNum) --> sliceSize
49 :
50 : SliceInfoVecforNHR: [1st chunk: [1st Slice, 2nd Slice, ...], 2nd chunk: [1st Slice, 2nd Slice, ...], ...]
51 : */
52 0 : HcclResult CcuTempAllToAllMesh1D::CalcSliceInfo(const AllignInfo &allignInfo, const u64 dataSize,
53 : RankSliceInfo &sliceInfoVec)
54 : {
55 0 : std::vector<SliceInfo> tmp(tempVTopo_.size());
56 0 : sliceInfoVec.resize(tempRankSize_, tmp);
57 :
58 0 : CHK_RET(CalcRsAgSliceInfoMesh(myRank_, tempRankSize_, allignInfo, dataSize, sliceInfoVec));
59 :
60 0 : return HcclResult::HCCL_SUCCESS;
61 0 : }
62 :
63 0 : void CcuTempAllToAllMesh1D::SetA2ASendRecvInfo(const A2ASendRecvInfo &sendRecvInfo)
64 : {
65 0 : localSendRecvInfo_ = sendRecvInfo;
66 0 : }
67 :
68 0 : HcclResult CcuTempAllToAllMesh1D::SetBuffBlockSize(const u64 buffBlockSize)
69 : {
70 0 : CHK_PRT_RET(buffBlockSize == 0, HCCL_ERROR("[CcuTempAllToAllMesh1D][SetBuffBlockSize] buffBlockSize should not be zero"),
71 : HcclResult::HCCL_E_PARA);
72 0 : buffBlockSize_ = buffBlockSize;
73 0 : return HcclResult::HCCL_SUCCESS;
74 : }
75 :
76 0 : HcclResult CcuTempAllToAllMesh1D::SetConcurrentSendRecvNum(const u32 concurrentSendRecvNum)
77 : {
78 0 : CHK_PRT_RET(concurrentSendRecvNum == 0, HCCL_ERROR("[CcuTempAllToAllMesh1D][SetConcurrentSendRecvNum] concurrentSendRecvNum should not be zero"),
79 : HcclResult::HCCL_E_PARA);
80 0 : concurrentSendRecvNum_ = concurrentSendRecvNum;
81 0 : return HcclResult::HCCL_SUCCESS;
82 : }
83 :
84 0 : uint64_t CcuTempAllToAllMesh1D::DataSliceToAddr(const DataSlice &dataSlice)
85 : {
86 0 : if (dataSlice.GetType() == BufferType::INPUT) {
87 0 : return static_cast<uint64_t>(op_.inputMem->GetAddr());
88 0 : } else if (dataSlice.GetType() == BufferType::OUTPUT) {
89 0 : return static_cast<uint64_t>(op_.outputMem->GetAddr());
90 : } else {
91 0 : return static_cast<uint64_t>(op_.scratchMem->GetAddr());
92 : }
93 : }
94 :
95 0 : HcclResult CcuTempAllToAllMesh1D::Run(const TempFuncs &tempFuncs, const RankSliceInfo &sliceInfoVec,
96 : const BuffInfo &buffInfo, const ResLinks &tempLinks,
97 : std::vector<InsQuePtr> &tempInsQues)
98 : {
99 0 : HCCL_INFO("[CcuTempAllToAllMesh1D] Run");
100 : (void)sliceInfoVec;
101 : (void)tempFuncs;
102 : (void)buffInfo;
103 0 : CcuInstructionAllToAllMesh1D ccuInsAllToAllMesh1D;
104 0 : CHK_PRT_RET(tempInsQues.empty(),
105 : HCCL_ERROR("[CcuTempAllToAllMesh1D] empty queue"), HcclResult::HCCL_E_INTERNAL);
106 0 : CHK_PTR_NULL(tempInsQues[0]);
107 0 : std::vector<uint64_t> dimSize;
108 0 : dimSize.push_back(tempRankSize_);
109 : // 拿到input和output的首地址,和每片小数据的大小
110 0 : uint64_t totalSliceSize = localSendRecvInfo_.sendLength[0]; // Bytes
111 0 : uint64_t inputAddr = op_.inputMem == nullptr ? 0 : static_cast<uint64_t>(op_.inputMem->GetAddr());
112 0 : uint64_t outputAddr = op_.outputMem == nullptr ? 0 : static_cast<uint64_t>(op_.outputMem->GetAddr());
113 : uint64_t token;
114 0 : CHK_RET(GetToken(op_, token));
115 0 : uint64_t srcStride = totalSliceSize + sendStrideSize_;
116 0 : uint64_t dstStride = totalSliceSize + recvStrideSize_;
117 :
118 0 : uint64_t loopCnt = totalSliceSize / UB_MAX_DATA_SIZE + ( totalSliceSize % UB_MAX_DATA_SIZE == 0 ? 0 : 1 );
119 0 : uint64_t sliceBias = 0;
120 0 : for ( uint64_t i = 0; i < loopCnt; i++ ) {
121 0 : if (tempRankSize_ == 1) {
122 : // ccu-alltoall算子的单P场景单独处理
123 0 : DataSlice usrInSlice = DataSlice(BufferType::INPUT, 0, totalSliceSize);
124 0 : DataSlice usrOutSlice = DataSlice(BufferType::OUTPUT, 0, totalSliceSize);
125 0 : std::unique_ptr<Instruction> insLocalCopy = std::make_unique<InsLocalCopy>(usrInSlice, usrOutSlice);
126 0 : tempInsQues[0]->Append(std::move(insLocalCopy));
127 0 : HCCL_INFO("[CcuTempAllToAllMesh1D] rankSize = 1, use InsLocalCopy for sliceSize[%llu].", totalSliceSize);
128 0 : break;
129 0 : }
130 : // prepare parameters & ccuIns init
131 0 : uint64_t sliceSize = ((i == loopCnt - 1) ? (totalSliceSize - i * UB_MAX_DATA_SIZE) : UB_MAX_DATA_SIZE);
132 0 : uint64_t srcOffset = sliceBias;
133 0 : uint64_t dstOffset = sliceBias + myRank_ * dstStride;
134 :
135 0 : ccuInsAllToAllMesh1D.Init(static_cast<uint32_t>(myRank_), inputAddr, outputAddr, sliceSize, token, srcOffset,
136 0 : dstOffset, srcStride, op_, tempVTopo_, loadFromMem_);
137 0 : HCCL_INFO("[CcuTempAllToAllMesh1D] Run Init: loadFromMem_[%d], myRank_[%d], dimSize[%llu], inputAddr[%llu],"\
138 : "outputAddr[%llu], sliceSize[%llu], srcOffset[%llu], dstOffset[%llu], loopCnt[%llu]",
139 : loadFromMem_, myRank_, dimSize[0], inputAddr, outputAddr, sliceSize, srcOffset, dstOffset, loopCnt);
140 : // init links
141 0 : std::vector<LinkData> links;
142 0 : for (auto &pair : tempLinks) {
143 0 : if (pair.second.empty()) {
144 0 : continue;
145 : }
146 0 : links.push_back(pair.second[0]);
147 : }
148 0 : HCCL_INFO("[CcuTempAllToAllMesh1D] links.size[%zu]", links.size());
149 0 : ccuInsAllToAllMesh1D.SetLinks(links);
150 0 : RankGroup rankGroup;
151 0 : for (auto &peer : tempVTopo_[0]) {
152 0 : rankGroup.AddRank(peer);
153 : }
154 0 : u32 cntCkeNum = 3;
155 0 : ccuInsAllToAllMesh1D.SetCntCkeNum(cntCkeNum);
156 0 : ccuInsAllToAllMesh1D.SetRankGroup(rankGroup);
157 0 : HCCL_INFO("CCUInsAllToAllmesh1D is [%s]", ccuInsAllToAllMesh1D.Describe().c_str());
158 0 : ccuInsAllToAllMesh1D.Describe();
159 0 : tempInsQues[0]->Append(std::move(std::make_unique<CcuInstructionAllToAllMesh1D>(ccuInsAllToAllMesh1D)));
160 0 : sliceBias += sliceSize;
161 0 : }
162 :
163 0 : return HcclResult::HCCL_SUCCESS;
164 0 : }
165 : } // namespace Hccl
|