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 CCU_INSTRUCTION_H
12 : #define CCU_INSTRUCTION_H
13 :
14 : #include <memory>
15 : #include "instruction.h"
16 : #include "ccu_ctx_arg.h"
17 : #include "ccu_task_arg.h"
18 : #include "ccu_rank_group.h"
19 : #include "ccu_ctx_signature.h"
20 : #include "ccu_task_param.h"
21 : #include "internal_exception.h"
22 :
23 : namespace Hccl {
24 2255 : MAKE_ENUM(CcuInstType, CCU_INS_GROUP, CCU_ALLTOALL_MESH_2D_DIRECT, CCU_ALLGATHER_MESH_1D_DIRECT,
25 : CCU_ALLGATHER_MESH_2D_DIRECT, CCU_REDUCE_SCATTER_MESH_1D_DIRECT, CCU_ALL_REDUCE_MESH_1D_DIRECT,
26 : CCU_ALLTOALL_MESH_1D_DIRECT, CCU_REDUCE_MESH_1D_DIRECT, CCU_ALLTOALLV_MESH_1D_DIRECT,
27 : CCU_SCATTER_MESH_1D_DIRECT, CCU_BROADCAST_MESH_1D_DIRECT, CCU_REDUCE_SCATTER_MESH_2D_DIRECT, CCU_REDUCE_SCATTER_MESH_2D_MULTI_MISSION,
28 : CCU_BROADCAST_MESH_2D_DIRECT, CCU_ALL_REDUCE_MESH_2D_ONE_SHOT_DIRECT, CCU_ALL_REDUCE_MESH_2D_TWO_SHOT_DIRECT,
29 : CCU_SCATTER_MESH_2D_DIRECT, CCU_REDUCE_MESH_2D_DIRECT, CCU_ALLTOALLV_MESH_2D_DIRECT,
30 : CCU_ALLGATHER_MESH_1D_DETOUR, CCU_ALL_REDUCE_MESH_1D_DETOUR, CCU_REDUCE_SCATTER_MESH_1D_DETOUR,
31 : CCU_BROADCAST_MESH_1D_MEM2MEM, CCU_BROADCAST_MESH_1D_MULTIMISSION, CCU_REDUCE_MESH_1D_MULTI_MISSION,
32 : CCU_ALL_REDUCE_MESH_1D_ONE_SHOT_DIRECT, CCU_REDUCE_TAILBLOCK_DIRECT, CCU_ALL_REDUCE_MESH_1D_MEM2MEM,
33 : CCU_ALL_REDUCE_MESH_1D_MULTI_MISSION, CCU_ALLGATHER_MESH_1D_MULTI_MISSION, CCU_ALLGATHER_MESH_1D_MEM2MEM,
34 : CCU_HALF_ALLTOALLV_MESH_1D, CCU_REDUCE_SCATTER_MESH_1D_MULTI_MISSION, CCU_REDUCE_SCATTER_MESH_1D_MEM2MEM, CCU_REDUCE_SCATTER_NHR_1D_MEM2MEM,
35 : CCU_ALLGATHER_MESH_1D_MEM2MEM_WITH_STRIDE_DIRECT, CCU_ALLGATHER_NHR_1D_MEM2MEM, CCU_ALLGATHER_MESH_2D_MULTI_MISSION,
36 : CCU_BROADCAST_MESH_2D_MEM2MEM, CCU_ALLGATHER_MESH_2D_MEM2MEM,
37 : CCU_ALL_GATHER_V_MESH_1D_DIRECT, CCU_REDUCE_SCATTER_V_MESH_1D_DIRECT, CCU_REDUCE_SCATTER_V_MESH_1D_MEM2MEM_DIRECT,
38 : CCU_ALL_REDUCE_MESH_2D_TWO_SHOT_MULTI_MISSION, CCU_ALL_REDUCE_MESH_2D_TWO_SHOT_MEM2MEM, CCU_REDUCE_MESH_1D_MEM2MEM,
39 : CCU_REDUCE_SCATTER_MESH_2D_MEM2MEM, CCU_BROADCAST_MESH_2D_MULTI_MISSION,
40 : CCU_REDUCE_MESH_2D_MEM2MEM, CCU_ALLREDUCE_NHR_1D_MEM2MEM,CCU_SCATTER_NHR_1D_MEM2MEM,CCU_BROADCAST_NHR_1D_MEM2MEM,
41 : CCU_REDUCE_NHR_1D_MEM2MEM, CCU_ALLTOALLV_MESH_2DIE_DIRECT, CCU_ALLGATHER_MESH_1D_2DIE, CCU_ALLTOALL_MESH_1D_2DIE, CCU_REDUCE_SCATTER_MESH_1D_2DIE,
42 : CCU_REDUCE_MESH_1D_TWO_SHOT_MEM2MEM);
43 :
44 : class CcuInstruction : public Instruction {
45 : public:
46 61 : CcuInstruction() : Instruction(InstructionType::CCU_INS)
47 : {
48 61 : }
49 :
50 3 : virtual void SetExecId(u64 id)
51 : {
52 3 : execId = id;
53 3 : }
54 :
55 3 : virtual u64 GetExecId() const
56 : {
57 3 : return execId;
58 : }
59 :
60 : u32 GetCntCkeNum() const
61 : {
62 : return cntCkeNum;
63 : }
64 :
65 4 : void SetCntCkeNum(u32 num)
66 : {
67 4 : cntCkeNum = num;
68 4 : }
69 :
70 6 : virtual CcuCtxSignature GetCtxSignature() const
71 : {
72 6 : std::unique_ptr<CcuCtxArg> ccuTaskArg = GetCtxArg();
73 6 : if (ccuTaskArg == nullptr) {
74 0 : THROW<InternalException>("[CcuInstruction][GetCtxSignature]ccuTaskArg is null");
75 : }
76 12 : return ccuTaskArg->GetCtxSignature();
77 6 : }
78 :
79 9 : virtual std::vector<LinkData> GetLinks() const
80 : {
81 9 : return links_;
82 : }
83 :
84 9 : virtual void SetLinks(std::vector<LinkData> &links)
85 : {
86 9 : links_ = links;
87 9 : }
88 :
89 0 : virtual RankGroup GetRankGroup() const
90 : {
91 0 : return rankGroup_;
92 : }
93 :
94 4 : virtual void SetRankGroup(RankGroup &rankGroup)
95 : {
96 4 : rankGroup_ = rankGroup;
97 4 : }
98 :
99 0 : virtual CcuInstType GetInstType() const
100 : {
101 0 : return instType_;
102 : }
103 :
104 : virtual void Translate(std::vector<std::vector<CcuTaskParam>> &taskParam) const;
105 :
106 : virtual std::unique_ptr<CcuCtxArg> GetCtxArg() const = 0;
107 : virtual std::unique_ptr<CcuTaskArg> GetTaskArg() const = 0;
108 : std::string Describe() const override = 0;
109 :
110 : protected:
111 : u64 execId{0};
112 : RankGroup rankGroup_;
113 : std::vector<LinkData> links_;
114 : CcuInstType instType_;
115 :
116 : private:
117 : u32 cntCkeNum{0};
118 : };
119 :
120 : } // namespace Hccl
121 :
122 : #endif // CCU_INSTRUCTION_H
|