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