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_DATA_TYPE_H
12 : #define HCCLV2_DATA_TYPE_H
13 :
14 : #include <map>
15 : #include <unordered_map>
16 : #include <string>
17 : #include <cstdint>
18 : #include "types.h"
19 : #include "../utils/enum_factory.h"
20 : #include "hccl_res.h"
21 : #include "log.h"
22 : #include "string_util.h"
23 : #include <hccl/hccl_types.h>
24 : #include "../utils/exception_util.h"
25 : #include "../exception/invalid_params_exception.h"
26 :
27 : namespace Hccl {
28 :
29 437445 : MAKE_ENUM(DataType, INT8, INT16, INT32, FP16, FP32, INT64, UINT64, UINT8, UINT16, UINT32,
30 : FP64, BFP16, INT128, BF16_SAT, HIF8, FP8E4M3, FP8E5M2, FP8E8M0)
31 :
32 : const std::unordered_map<DataType, u32, std::EnumClassHash> DATA_TYPE_SIZE_MAP = {
33 : {DataType::INT8, sizeof(s8)},
34 : {DataType::INT16, sizeof(s16)},
35 : {DataType::INT32, sizeof(s32)},
36 : {DataType::FP16, 2},
37 : {DataType::FP32, sizeof(float)},
38 : {DataType::INT64, sizeof(s64)},
39 : {DataType::UINT64, sizeof(u64)},
40 : {DataType::UINT8, sizeof(u8)},
41 : {DataType::UINT16, sizeof(u16)},
42 : {DataType::UINT32, sizeof(u32)},
43 : {DataType::FP64, 8},
44 : {DataType::BFP16, 2},
45 : {DataType::INT128, 16},
46 : {DataType::BF16_SAT, 2},
47 : {DataType::HIF8, 1},
48 : {DataType::FP8E4M3, 1},
49 : {DataType::FP8E5M2, 1},
50 : {DataType::FP8E8M0, 1}
51 : };
52 :
53 : const std::unordered_map<DataType, HcclDataType, std::EnumClassHash> HCCL_DATA_TYPE_MAP = {
54 : {DataType::INT8, HCCL_DATA_TYPE_INT8},
55 : {DataType::INT16, HCCL_DATA_TYPE_INT16},
56 : {DataType::INT32, HCCL_DATA_TYPE_INT32},
57 : {DataType::FP16, HCCL_DATA_TYPE_FP16},
58 : {DataType::FP32, HCCL_DATA_TYPE_FP32},
59 : {DataType::INT64, HCCL_DATA_TYPE_INT64},
60 : {DataType::UINT64, HCCL_DATA_TYPE_UINT64},
61 : {DataType::UINT8, HCCL_DATA_TYPE_UINT8},
62 : {DataType::UINT16, HCCL_DATA_TYPE_UINT16},
63 : {DataType::UINT32, HCCL_DATA_TYPE_UINT32},
64 : {DataType::FP64, HCCL_DATA_TYPE_FP64},
65 : {DataType::BFP16, HCCL_DATA_TYPE_BFP16},
66 : {DataType::INT128, HCCL_DATA_TYPE_INT128},
67 : {DataType::HIF8, HCCL_DATA_TYPE_HIF8},
68 : {DataType::FP8E4M3, HCCL_DATA_TYPE_FP8E4M3},
69 : {DataType::FP8E5M2, HCCL_DATA_TYPE_FP8E5M2},
70 : {DataType::FP8E8M0, HCCL_DATA_TYPE_FP8E8M0},
71 : };
72 :
73 : const std::unordered_map<HcclDataType, DataType, std::EnumClassHash> DATA_TYPE_MAP = {
74 : {HCCL_DATA_TYPE_INT8, DataType::INT8},
75 : {HCCL_DATA_TYPE_INT16, DataType::INT16},
76 : {HCCL_DATA_TYPE_INT32, DataType::INT32},
77 : {HCCL_DATA_TYPE_FP16, DataType::FP16},
78 : {HCCL_DATA_TYPE_FP32, DataType::FP32},
79 : {HCCL_DATA_TYPE_INT64, DataType::INT64},
80 : {HCCL_DATA_TYPE_UINT64, DataType::UINT64},
81 : {HCCL_DATA_TYPE_UINT8, DataType::UINT8},
82 : {HCCL_DATA_TYPE_UINT16, DataType::UINT16},
83 : {HCCL_DATA_TYPE_UINT32, DataType::UINT32},
84 : {HCCL_DATA_TYPE_FP64, DataType::FP64},
85 : {HCCL_DATA_TYPE_BFP16, DataType::BFP16},
86 : {HCCL_DATA_TYPE_INT128, DataType::INT128},
87 : {HCCL_DATA_TYPE_HIF8, DataType::HIF8},
88 : {HCCL_DATA_TYPE_FP8E4M3, DataType::FP8E4M3},
89 : {HCCL_DATA_TYPE_FP8E5M2, DataType::FP8E5M2},
90 : {HCCL_DATA_TYPE_FP8E8M0, DataType::FP8E8M0},
91 : };
92 :
93 : const std::unordered_map<uint32_t, std::string> DATA_TYPE_TO_STRING_MAP = {
94 : {0, "INT8"},
95 : {1, "INT16"},
96 : {2, "INT32"},
97 : {3, "FP16"},
98 : {4, "FP32"},
99 : {5, "INT64"},
100 : {6, "UINT64"},
101 : {7, "UINT8"},
102 : {8, "UINT16"},
103 : {9, "UINT32"},
104 : {10, "FP64"},
105 : {11, "BFP16"},
106 : {12, "INT128"},
107 : {14, "HIF8"},
108 : {15, "FP8E4M3"},
109 : {16, "FP8E5M2"},
110 : {17, "FP8E8M0"},
111 : };
112 :
113 : const std::unordered_map<uint32_t, std::string> OP_TYPE_TO_STRING_MAP = {{0, "sum"}, {1, "mul"}, {2, "max"}, {3, "PROD"}};
114 :
115 139 : inline u32 DataTypeSizeGet(DataType type)
116 : {
117 139 : if (UNLIKELY(DATA_TYPE_SIZE_MAP.find(type) == DATA_TYPE_SIZE_MAP.end())) {
118 0 : THROW<InvalidParamsException>(StringFormat("%s type[%s] is not supported.", __func__, type.Describe().c_str()));
119 : }
120 139 : return DATA_TYPE_SIZE_MAP.at(type);
121 : }
122 :
123 8 : inline HcclDataType DataTypeToHcclDataType(const DataType dataType)
124 : {
125 8 : if (UNLIKELY(HCCL_DATA_TYPE_MAP.find(dataType) == HCCL_DATA_TYPE_MAP.end())) {
126 0 : THROW<InvalidParamsException>(StringFormat("%s type[%s] is not supported.", __func__, dataType.Describe().c_str()));
127 : }
128 8 : return HCCL_DATA_TYPE_MAP.at(dataType);
129 : }
130 :
131 199 : inline DataType HcclDataTypeToDataType(const HcclDataType hcclDataType)
132 : {
133 199 : if (UNLIKELY(DATA_TYPE_MAP.find(hcclDataType) == DATA_TYPE_MAP.end())) {
134 3 : HCCL_ERROR("%s hcclDataType[%d] is not supported.", __func__, hcclDataType);
135 1 : return DataType::INVALID;
136 : }
137 198 : return DATA_TYPE_MAP.at(hcclDataType);
138 : }
139 :
140 21 : inline std::string DataTypeToSerialString(const uint32_t dataType)
141 : {
142 21 : if (UNLIKELY(DATA_TYPE_TO_STRING_MAP.find(dataType) == DATA_TYPE_TO_STRING_MAP.end())) {
143 42 : return "UNDEFINED";
144 : }
145 0 : return DATA_TYPE_TO_STRING_MAP.at(dataType);
146 : }
147 :
148 1 : inline std::string OpTypeToSerialString(const uint32_t opType)
149 : {
150 1 : if (OP_TYPE_TO_STRING_MAP.find(opType) == OP_TYPE_TO_STRING_MAP.end()) {
151 0 : return "UNDEFINED";
152 : }
153 1 : return OP_TYPE_TO_STRING_MAP.at(opType);
154 : }
155 :
156 : } // namespace Hccl
157 :
158 : #endif // HCCLV2_DATA_TYPE_H
|