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 : #ifndef HCCLV2_INS_TEMP_SCATTER_MESH_2D
12 : #define HCCLV2_INS_TEMP_SCATTER_MESH_2D
13 :
14 : #include "string_util.h"
15 :
16 : #include "ins_alg_template_base.h"
17 : #include "executor_utils.h"
18 :
19 : namespace Hccl {
20 :
21 : class InsTempScatterMesh2D : public InsAlgTemplateBase {
22 : public:
23 : explicit InsTempScatterMesh2D(const RankId virtualRank, const u32 tempRankSize,
24 : const std::vector<std::vector<RankId>> &tempVTopo,
25 : const std::map<RankId, u32> &tempVirtRankMap);
26 : ~InsTempScatterMesh2D() override;
27 :
28 0 : std::string Describe() const override
29 : {
30 0 : return StringFormat("Instruction based Template of scatter mesh 2D with tempRankSize [%u].", tempRankSize_);
31 : }
32 :
33 : HcclResult GenExtIns(TempFuncs &tempFuncs, TemplateDataParams &tempAlgParams,
34 : ResLinks &tempResLinks, std::vector<InsQuePtr> &tempInsQues);
35 : HcclResult CalcRes(AlgTempResReq &tempResReq) override;
36 : u32 CalcScratchMultiple(BufferType inBuffType, BufferType outBuffType) const;
37 :
38 : private:
39 : HcclResult PreDataCopy(const TemplateDataParams &tempAlgParams, std::vector<InsQuePtr> &tempInsQues) const;
40 : HcclResult PostDataCopy(const TemplateDataParams &tempAlgParams, std::vector<InsQuePtr> &tempInsQues) const;
41 :
42 : HcclResult RootMeshSend(TemplateDataParams &tempAlgParams, ResLinks &tempResLinks,
43 : const std::vector<RankId> vTopo, const u32 xyRankSize, u32 rankDistX, u32 rankDistY, u64 dataOffSet,
44 : u64 tranDataSize, std::vector<InsQuePtr> &tempInsQues) const;
45 : HcclResult RankRecvFromRoot(TemplateDataParams &tempAlgParams, ResLinks &tempResLinks,
46 : const u32 xyRankSize, u32 rankDist, u64 dataOffSet, u64 tranDataSize, std::vector<InsQuePtr> &tempInsQues) const;
47 : HcclResult RunDataCombine(TemplateDataParams &tempAlgParams, ResLinks &tempResLinks, u64 tranDataSizeX,
48 : u64 tranDataSizeY, std::vector<InsQuePtr> &tempInsQues);
49 :
50 : HcclResult DataCombineSend(const TemplateDataParams &tempAlgParams, ResLinks &tempResLinks,
51 : const std::vector<RankId> vTopo, u64 dataOffSet, u64 tranDataSize,
52 : std::vector<InsQuePtr> &tempInsQues) const;
53 : HcclResult DataCombineRecv(const TemplateDataParams &tempAlgParams, LinkData link, u64 dataOffSet,
54 : u64 tranDataSize, InsQuePtr queue) const;
55 :
56 : HcclResult GetRankId(u32 xRank, u32 yRank, u32 &rank) const;
57 :
58 : HcclResult SendDirect(TemplateDataParams &tempAlgParams, InsQuePtr queue,
59 : const LinkData link, u32 remoteRank) const;
60 : HcclResult SendTransit(const TemplateDataParams &tempAlgParams, InsQuePtr queue,
61 : const LinkData link, u32 remoteRank, u32 xyRankSize, u32 rankDistX, u32 rankDistY, u64 xyOffSet, u64 tranDataSize) const;
62 : HcclResult RecvDirect(TemplateDataParams &tempAlgParams, InsQuePtr queue, const LinkData link, u32 remoteRank) const;
63 : HcclResult RecvTransit(const TemplateDataParams &tempAlgParams, InsQuePtr queue, const LinkData link,
64 : u32 remoteRank, u32 xyRankSize, u32 rankDist, u64 xyOffSet, u64 tranDataSize) const;
65 :
66 : u32 queNum_{0};
67 : u32 xQueNum_{0};
68 : u32 yQueNum_{0};
69 :
70 : u32 xRankSize_{0};
71 : u32 yRankSize_{0};
72 :
73 : u32 myRankX_{0};
74 : u32 myRankY_{0};
75 : u32 rootX_{0};
76 : u32 rootY_{0};
77 :
78 : std::vector<InsQuePtr> xInsQues_;
79 : std::vector<InsQuePtr> yInsQues_;
80 :
81 : bool enableInterRankCounterNotify_{false};
82 : };
83 :
84 : } // namespace Hccl
85 :
86 : #endif // HCCLV2_INS_TEMP_SCATTER_MESH_2D
|