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