已合并
[A2/A3]新增TransposeBatchMatMul算子NZ格式aclnn接口 #1702
[A2/A3]新增TransposeBatchMatMul算子NZ格式aclnn接口 #1702
已合并
fbbccc创建于 2月9日
共 5 个文件变更+796-19
@@ -36,11 +36,13 @@ using namespace op;
36using namespace Ops::NN;36using namespace Ops::NN;
37static const std::initializer_list<op::DataType> x1_SUPPORT_LIST = {DataType::DT_FLOAT, DataType::DT_FLOAT16, DataType::DT_BF16};37static const std::initializer_list<op::DataType> x1_SUPPORT_LIST = {DataType::DT_FLOAT, DataType::DT_FLOAT16, DataType::DT_BF16};
38static const std::initializer_list<op::DataType> x2_SUPPORT_LIST = {DataType::DT_FLOAT, DataType::DT_FLOAT16, DataType::DT_BF16};38static const std::initializer_list<op::DataType> x2_SUPPORT_LIST = {DataType::DT_FLOAT, DataType::DT_FLOAT16, DataType::DT_BF16};
39+static const std::initializer_list<op::DataType> DTYPE_SUPPORT_LIST_WEIGHTNZ = {op::DataType::DT_FLOAT16, op::DataType::DT_BF16};
39static const std::initializer_list<op::DataType> x1_SCALE_SUPPORT_LIST = {DataType::DT_FLOAT16};40static const std::initializer_list<op::DataType> x1_SCALE_SUPPORT_LIST = {DataType::DT_FLOAT16};
40static const std::initializer_list<op::DataType> x2_SCALE_SUPPORT_LIST = {DataType::DT_FLOAT16};41static const std::initializer_list<op::DataType> x2_SCALE_SUPPORT_LIST = {DataType::DT_FLOAT16};
41static const std::initializer_list<op::DataType> SCALE_DTYPE_SUPPORT_LIST = {DataType::DT_INT64, DataType::DT_UINT64};42static const std::initializer_list<op::DataType> SCALE_DTYPE_SUPPORT_LIST = {DataType::DT_INT64, DataType::DT_UINT64};
42static const std::initializer_list<op::DataType> OUT_DTYPE_SUPPORT_LIST = {DataType::DT_FLOAT, DataType::DT_FLOAT16, DataType::DT_BF16, DataType::DT_INT8};43static const std::initializer_list<op::DataType> OUT_DTYPE_SUPPORT_LIST = {DataType::DT_FLOAT, DataType::DT_FLOAT16, DataType::DT_BF16, DataType::DT_INT8};
43static constexpr size_t EXPECTED_DIM = 3;44static constexpr size_t EXPECTED_DIM = 3;
45+static constexpr int EXPECTED_NZ_DIM = 5;
44static constexpr int BLOCK_SIZE = 16;46static constexpr int BLOCK_SIZE = 16;
45static constexpr int SUPPORTED_INNER_AXIS = 65536;47static constexpr int SUPPORTED_INNER_AXIS = 65536;
46 48 
@@ -59,6 +61,43 @@ inline static bool CheckNotNull(const aclTensor* x1, const aclTensor* x2, const
59inline static bool CheckDtypeValid(const aclTensor* x1, const aclTensor* x2,61inline static bool CheckDtypeValid(const aclTensor* x1, const aclTensor* x2,
sxb154714
sxb154714sxb1547142月26日

在这个分支增加分支

likedislike
60 const aclTensor* scale, const aclTensor* out)62 const aclTensor* scale, const aclTensor* out)
61{63{
64+ if (x1->GetDataType() != x2->GetDataType()) {
65+ OP_LOGE(
66+ ACLNN_ERR_PARAM_INVALID,
67+ "x1's dtype [%s] and x2's dtype [%s] are not equal.",
68+ op::ToString(x1->GetDataType()).GetString(), op::ToString(x2->GetDataType()).GetString());
69+ return false;
70+ }
71+ // Handle weight NZ format specific checks
72+ if (x2->GetStorageFormat() == Format::FORMAT_FRACTAL_NZ) {
73+ auto socVersion = GetCurrentPlatformInfo().GetSocVersion();
74+ auto npuArch = GetCurrentPlatformInfo().GetCurNpuArch();
75+ if (npuArch != NpuArch::DAV_2201) {
76+ OP_LOGE(
77+ ACLNN_ERR_PARAM_INVALID,
78+ "transposebatchmatmulweightnz is unsupported by the current SOC version [%s].",
79+ op::ToString(socVersion).GetString());
80+ return false;
81+ }
82+ OP_CHECK_DTYPE_NOT_SUPPORT(x1, DTYPE_SUPPORT_LIST_WEIGHTNZ, return false);
83+ OP_CHECK_DTYPE_NOT_SUPPORT(x2, DTYPE_SUPPORT_LIST_WEIGHTNZ, return false);
84+ OP_CHECK_DTYPE_NOT_SUPPORT(out, DTYPE_SUPPORT_LIST_WEIGHTNZ, return false);
85+ if (scale != nullptr) {
86+ OP_CHECK_DTYPE_NOT_SUPPORT(scale, SCALE_DTYPE_SUPPORT_LIST, return false);
87+ OP_CHECK_DTYPE_NOT_SUPPORT(x1, x1_SCALE_SUPPORT_LIST, return false);
88+ OP_CHECK_DTYPE_NOT_SUPPORT(x2, x2_SCALE_SUPPORT_LIST, return false);
89+ }
90+ if (x1->GetDataType() != out->GetDataType()) {
91+ OP_LOGE(
92+ ACLNN_ERR_PARAM_INVALID,
93+ "x1's dtype [%s] and out's dtype [%s] are not equal.",
94+ op::ToString(x1->GetDataType()).GetString(), op::ToString(out->GetDataType()).GetString());
95+ return false;
96+ }
97+ return true;
98+ }
99+
100+ // Regular ND format checks
62 OP_CHECK_DTYPE_NOT_SUPPORT(x1, x1_SUPPORT_LIST, return false);101 OP_CHECK_DTYPE_NOT_SUPPORT(x1, x1_SUPPORT_LIST, return false);
63 OP_CHECK_DTYPE_NOT_SUPPORT(x2, x2_SUPPORT_LIST, return false);102 OP_CHECK_DTYPE_NOT_SUPPORT(x2, x2_SUPPORT_LIST, return false);
64 if (scale != nullptr) {103 if (scale != nullptr) {
@@ -172,6 +211,17 @@ static inline bool CheckMathType(const aclTensor* x1, const aclTensor* x2, int8_
172 return CheckCubeMathTypeForMm(promoteType, cubeMathType);211 return CheckCubeMathTypeForMm(promoteType, cubeMathType);
173}212}
174 213 
214+static inline bool CheckNzStorageShape(const aclTensor* x2)
215+{
216+ auto storageShape = x2->GetStorageShape();
217+ auto storageShapeDim = storageShape.GetDimNum();
218+ OP_CHECK(
219+ storageShapeDim == EXPECTED_NZ_DIM,
220+ OP_LOGE(ACLNN_ERR_PARAM_INVALID, "Only support x2 storageShapeDim is 5, which are [%zu].", storageShapeDim),
221+ return false);
222+ return true;
223+}
224+ 
175inline static aclnnStatus CheckParams(const aclTensor* x1, const aclTensor* x2, const aclTensor* scale, aclTensor* out,225inline static aclnnStatus CheckParams(const aclTensor* x1, const aclTensor* x2, const aclTensor* scale, aclTensor* out,
176 const aclIntArray* perm_x1, const aclIntArray* perm_x2, const aclIntArray* perm_y,226 const aclIntArray* perm_x1, const aclIntArray* perm_x2, const aclIntArray* perm_y,
177 int8_t cubeMathType, int32_t batch_split_factor)227 int8_t cubeMathType, int32_t batch_split_factor)
@@ -206,6 +256,9 @@ inline static aclnnStatus CheckParams(const aclTensor* x1, const aclTensor* x2,
206 256 
207 CHECK_RET(CheckDtypeValid(x1, x2, scale, out), ACLNN_ERR_PARAM_INVALID);257 CHECK_RET(CheckDtypeValid(x1, x2, scale, out), ACLNN_ERR_PARAM_INVALID);
208 CHECK_RET(CheckShapeValid(x1, x2, scale, perm_x1, perm_x2), ACLNN_ERR_PARAM_INVALID);258 CHECK_RET(CheckShapeValid(x1, x2, scale, perm_x1, perm_x2), ACLNN_ERR_PARAM_INVALID);
259+ if (x2->GetStorageFormat() == Format::FORMAT_FRACTAL_NZ) {
260+ CHECK_RET(CheckNzStorageShape(x2), ACLNN_ERR_PARAM_INVALID);
261+ }
209 return ACLNN_SUCCESS;262 return ACLNN_SUCCESS;
210}263}
211 264 
@@ -239,7 +292,7 @@ static const aclTensor* BuildTransposeBatchMatMulGraph(const aclTensor* x1, cons
239 if (contiguousScale != nullptr) {292 if (contiguousScale != nullptr) {
240 contiguousScale = l0op::Contiguous(scale, executor);293 contiguousScale = l0op::Contiguous(scale, executor);
241 OP_CHECK(contiguousScale != nullptr, OP_LOGE(ACLNN_ERR_INNER_NULLPTR,294 OP_CHECK(contiguousScale != nullptr, OP_LOGE(ACLNN_ERR_INNER_NULLPTR,
242- "THe input scale perprocess failed, contiguouse return nullptr."),295+ "The input scale perprocess failed, contiguouse return nullptr."),
243 return nullptr);296 return nullptr);
244 }297 }
245 298 
@@ -248,6 +301,47 @@ static const aclTensor* BuildTransposeBatchMatMulGraph(const aclTensor* x1, cons
248 perm_y, cubeMathType == USE_HF32, batch_split_factor, executor);301 perm_y, cubeMathType == USE_HF32, batch_split_factor, executor);
249}302}
250 303 
304+static const aclTensor* BuildTransposeBatchMatMulWeightNzGraph(const aclTensor* x1, const aclTensor* x2,
305+ const aclTensor* scale, const aclIntArray* perm_x1,
306+ const aclIntArray* perm_x2, const aclIntArray* perm_y,
307+ int8_t cubeMathType, int32_t batch_split_factor,
308+ aclOpExecutor *executor)
309+{
310+ // 连续性转换
311+ auto contiguousX1 = l0op::Contiguous(x1, executor);
312+ OP_CHECK(contiguousX1 != nullptr, OP_LOGE(ACLNN_ERR_INNER_NULLPTR,
313+ "The input x1 perprocess failed, contiguouse return nullptr."),
314+ return nullptr);
315+ auto reformX1 = l0op::ReFormat(contiguousX1, op::Format::FORMAT_ND);
316+ OP_CHECK(reformX1 != nullptr, OP_LOGE(ACLNN_ERR_INNER_NULLPTR,
317+ "The input x1 perprocess failed, reformat return nullptr."),
318+ return nullptr);
319+ 
320+ // 原始方法传入
321+ auto contiguousX2 = l0op::Contiguous(x2, executor);
322+ OP_CHECK(contiguousX2 != nullptr, OP_LOGE(ACLNN_ERR_INNER_NULLPTR,
323+ "The input x2 perprocess failed, contiguouse return nullptr."),
324+ return nullptr);
325+ 
326+ // weightnz storageshape刷新
327+ if (x2->GetStorageFormat() == op::Format::FORMAT_FRACTAL_NZ) {
328+ contiguousX2->SetStorageShape(x2->GetStorageShape());
329+ }
330+ 
331+ // scale非连续转连续以及转换dtype
332+ auto contiguousScale = scale;
333+ if (contiguousScale != nullptr) {
334+ contiguousScale = l0op::Contiguous(scale, executor);
335+ OP_CHECK(contiguousScale != nullptr, OP_LOGE(ACLNN_ERR_INNER_NULLPTR,
336+ "The input scale perprocess failed, contiguouse return nullptr."),
337+ return nullptr);
338+ }
339+ 
340+ // 构建matmul计算图
341+ return l0op::TransposeBatchMatMul(reformX1, contiguousX2, nullptr, contiguousScale, perm_x1, perm_x2,
342+ perm_y, cubeMathType == USE_HF32, batch_split_factor, executor);
343+}
344+ 
251aclnnStatus aclnnTransposeBatchMatMulGetWorkspaceSize(const aclTensor* x1, const aclTensor* x2, const aclTensor* bias,345aclnnStatus aclnnTransposeBatchMatMulGetWorkspaceSize(const aclTensor* x1, const aclTensor* x2, const aclTensor* bias,
252 const aclTensor* scale, const aclIntArray* permX1,346 const aclTensor* scale, const aclIntArray* permX1,
253 const aclIntArray* permX2, const aclIntArray* permY,347 const aclIntArray* permX2, const aclIntArray* permY,
@@ -303,6 +397,71 @@ aclnnStatus aclnnTransposeBatchMatMul(void* workspace, uint64_t workspaceSize, a
303 const aclrtStream stream)397 const aclrtStream stream)
304{398{
305 L2_DFX_PHASE_2(aclnnTransposeBatchMatMul);399 L2_DFX_PHASE_2(aclnnTransposeBatchMatMul);
306- 400+ return CommonOpExecutorRun(workspace, workspaceSize, executor, stream);
401+}
402+ 
403+aclnnStatus aclnnTransposeBatchMatMulWeightNzGetWorkspaceSize(const aclTensor* x1, const aclTensor* x2, const aclTensor* bias,
404+ const aclTensor* scale, const aclIntArray* permX1,
405+ const aclIntArray* permX2, const aclIntArray* permY,
406+ int8_t cubeMathType, const int32_t batchSplitFactor,
407+ aclTensor* out, uint64_t* workspaceSize,
408+ aclOpExecutor** executor)
409+{
410+ L2_DFX_PHASE_1(aclnnTransposeBatchMatMulWeightNz,
411+ DFX_IN(x1, x2, bias, scale, permX1, permX2, permY, cubeMathType, batchSplitFactor), DFX_OUT(out));
412+ 
413+ // 固定写法, 创建OpExecutor
414+ auto unique_executor = CREATE_EXECUTOR();
415+ CHECK_RET(unique_executor.get() != nullptr, ACLNN_ERR_INNER_CREATE_EXECUTOR);
416+ 
417+ // x2 format must be NZ
林旭
林旭林旭3月3日

这段校验是否可以放到checkParams中,或者封装一个checkParamsWeightNz复用checkParams

likedislike
418+ if (ge::GetPrimaryFormat(x2->GetStorageFormat()) != Format::FORMAT_FRACTAL_NZ) {
419+ OP_LOGE(
420+ ACLNN_ERR_PARAM_INVALID, "Format of x2 must be FRACTAL_NZ, actual is %s.",
421+ op::ToString(x2->GetStorageFormat()).GetString());
422+ return ACLNN_ERR_PARAM_INVALID;
423+ }
424+ 
425+ // 入参检查
426+ auto ret = CheckParams(x1, x2, scale, out, permX1, permX2, permY, cubeMathType, batchSplitFactor);
427+ CHECK_RET(ret == ACLNN_SUCCESS, ret);
428+ 
429+ // 空tensor 处理
430+ if (x1->IsEmpty() || x2->IsEmpty()) {
431+ OP_LOGE(ACLNN_ERR_PARAM_INVALID, "aclnnTransposeBatchMatMulWeightNz do not support empty tensor!");
432+ return ACLNN_ERR_PARAM_INVALID;
433+ }
434+ 
435+ if (bias != nullptr) {
436+ OP_LOGE(ACLNN_ERR_PARAM_INVALID, "The bias is not support in TBMM.");
437+ return ACLNN_ERR_PARAM_INVALID;
438+ }
439+ 
440+ // 构建matmul计算图
441+ const aclTensor* tbmmOut = nullptr;
442+ tbmmOut = BuildTransposeBatchMatMulWeightNzGraph(x1, x2, scale, permX1, permX2, permY,
443+ cubeMathType, batchSplitFactor, unique_executor.get());
444+ CHECK_RET(tbmmOut != nullptr, ACLNN_ERR_PARAM_INVALID);
445+ 
446+ if (tbmmOut->IsEmpty()) {
447+ *workspaceSize = 0;
448+ unique_executor.ReleaseTo(executor);
449+ return ACLNN_SUCCESS;
450+ }
451+ 
452+ tbmmOut = l0op::Cast(tbmmOut, out->GetDataType(), unique_executor.get());
453+ CHECK_RET(tbmmOut != nullptr, ACLNN_ERR_PARAM_INVALID);
454+ auto viewCopyResult = l0op::ViewCopy(tbmmOut, out, unique_executor.get());
455+ CHECK_RET(viewCopyResult != nullptr, ACLNN_ERR_PARAM_INVALID);
456+ 
457+ *workspaceSize = unique_executor->GetWorkspaceSize();
458+ unique_executor.ReleaseTo(executor);
459+ return ACLNN_SUCCESS;
460+}
461+ 
462+aclnnStatus aclnnTransposeBatchMatMulWeightNz(void* workspace, uint64_t workspaceSize, aclOpExecutor* executor,
463+ const aclrtStream stream)
464+{
465+ L2_DFX_PHASE_2(aclnnTransposeBatchMatMulWeightNz);
307 return CommonOpExecutorRun(workspace, workspaceSize, executor, stream);466 return CommonOpExecutorRun(workspace, workspaceSize, executor, stream);
308}467}
@@ -49,6 +49,42 @@ ACLNN_API aclnnStatus aclnnTransposeBatchMatMulGetWorkspaceSize(const aclTensor*
49ACLNN_API aclnnStatus aclnnTransposeBatchMatMul(void* workspace, uint64_t workspaceSize, aclOpExecutor* executor,49ACLNN_API aclnnStatus aclnnTransposeBatchMatMul(void* workspace, uint64_t workspaceSize, aclOpExecutor* executor,
50 const aclrtStream stream);50 const aclrtStream stream);
51 51 
52+/**
53+ * @brief aclnnTransposeBatchMatMulWeightNz的第一段接口,根据具体的计算流程,计算workspace大小。
54+ * @domain aclnn_ops_infer
55+ * 算子功能:相对于aclnnTransposeBatchMatmul, mat2为NZ格式。
56+ * @param [in] x1: matmul左矩阵,数据类型支持:float16、bfloat16。数据格式支持ND。
57+ * @param [in] x2: matmul右矩阵,数据类型支持:float16、bfloat16。数据格式支持NZ。
58+ * @param [in] bias: 偏置,当前不支持。
59+ * @param [in] scale: 量化参数中的缩放因子,数据类型支持:int64、uint64。
60+ * @param [in] permX1: 表示输入x1的shape。
61+ * @param [in] permX2: 表示输入x2的shape。
62+ * @param [in] permY: 表示输入y的shape。
63+ * @param [in] cubeMathType: 用于指定Cube单元的计算逻辑,Host侧的整型。数据类型支持:int8。
64+ * @param [in] batchSplitFactor: 是否重新拆分shape。数据类型支持:int32。
65+ * @param [out] out: 计算结果,数据类型:float16, bfloat16。数据格式支持ND。
66+ * @param [out] workspaceSize: 返回需要在npu device侧申请的workspace大小。
67+ * @param [out] executor: 返回op执行器,包含了算子计算流程。
68+ * @return aclnnStatus: 返回状态码
69+ */
70+ACLNN_API aclnnStatus aclnnTransposeBatchMatMulWeightNzGetWorkspaceSize(const aclTensor* x1, const aclTensor* x2,
71+ const aclTensor* bias, const aclTensor* scale,
72+ const aclIntArray* permX1, const aclIntArray* permX2,
73+ const aclIntArray* permY, int8_t cubeMathType,
74+ const int32_t batchSplitFactor, aclTensor* out,
75+ uint64_t* workspaceSize, aclOpExecutor** executor);
76+ 
77+/**
78+ * @brief aclnnTransposeBatchMatMulWeightNz的第二段接口,用于执行计算。
79+ * @param [in] workspace: 在npu device侧申请的workspace内存起址。
80+ * @param [in] workspace_size: 在npu device侧申请的workspace大小,由第一段接口aclnnBatchMatMulWeightNzGetWorkspaceSize获取。
81+ * @param [in] executor: op执行器,包含了算子计算流程。
82+ * @param [in] stream: acl stream流。
83+ * @return aclnnStatus: 返回状态码。
84+ */
85+ACLNN_API aclnnStatus aclnnTransposeBatchMatMulWeightNz(void* workspace, uint64_t workspaceSize,
86+ aclOpExecutor* executor, const aclrtStream stream);
87+ 
52#ifdef __cplusplus88#ifdef __cplusplus
53}89}
54#endif90#endif
@@ -2,9 +2,203 @@
2 "op_type": "TransposeBatchMatMul",2 "op_type": "TransposeBatchMatMul",
3 "optional_input_mode": "gen_placeholder",3 "optional_input_mode": "gen_placeholder",
4 "op_list": [4 "op_list": [
5+ {
6+ "bin_filename": "TransposeBatchMatMul_ND_NZ_ND_ND_ND_FP16_FP16_FP32_INT64_FP16",
7+ "simplified_key": "diy,2/29/2/2/2/1/1/0/9/1",
8+ "inputs": [
9+ {
10+ "name": "x1",
11+ "index": 0,
12+ "dtype": "float16",
13+ "format": "ND",
14+ "paramType": "required",
15+ "shape": [
16+ -2
17+ ]
18+ },
19+ {
20+ "name": "x2",
21+ "index": 1,
22+ "dtype": "float16",
23+ "format": "FRACTAL_NZ",
24+ "paramType": "required",
25+ "shape": [
26+ -2
27+ ]
28+ },
29+ {
30+ "name": "bias",
31+ "index": 2,
32+ "dtype": "float32",
33+ "format": "ND",
34+ "paramType": "optional",
35+ "shape": [
36+ -2
37+ ]
38+ },
39+ {
40+ "name": "scale",
41+ "index": 3,
42+ "dtype": "int64",
43+ "format": "ND",
44+ "paramType": "optional",
45+ "shape": [
46+ -2
47+ ]
48+ }
49+ ],
50+ "outputs": [
51+ {
52+ "name": "y",
53+ "index": 0,
54+ "dtype": "float16",
55+ "format": "ND",
56+ "paramType": "required",
57+ "shape": [
58+ -2
59+ ]
60+ }
61+ ],
62+ "attrs": [
63+ {
64+ "name": "perm_x1",
65+ "dtype": "list_int",
66+ "value": [
67+ 1,
68+ 0,
69+ 2
70+ ]
71+ },
72+ {
73+ "name": "perm_x2",
74+ "dtype": "list_int",
75+ "value": [
76+ 0,
77+ 1,
78+ 2
79+ ]
80+ },
81+ {
82+ "name": "perm_y",
83+ "dtype": "list_int",
84+ "value": [
85+ 1,
86+ 0,
87+ 2
88+ ]
89+ },
90+ {
91+ "name": "enable_hf32",
92+ "dtype": "bool",
93+ "value": false
94+ },
95+ {
96+ "name": "batch_split_factor",
97+ "dtype": "int",
98+ "value": 1
99+ }
100+ ]
101+ },
102+ {
103+ "bin_filename": "TransposeBatchMatMul_ND_NZ_ND_ND_ND_FP16_FP16_FP16_INT64_FP16",
104+ "simplified_key": "diy,2/29/2/2/2/1/1/1/9/1",
105+ "inputs": [
106+ {
107+ "name": "x1",
108+ "index": 0,
109+ "dtype": "float16",
110+ "format": "ND",
111+ "paramType": "required",
112+ "shape": [
113+ -2
114+ ]
115+ },
116+ {
117+ "name": "x2",
118+ "index": 1,
119+ "dtype": "float16",
120+ "format": "FRACTAL_NZ",
121+ "paramType": "required",
122+ "shape": [
123+ -2
124+ ]
125+ },
126+ {
127+ "name": "bias",
128+ "index": 2,
129+ "dtype": "float16",
130+ "format": "ND",
131+ "paramType": "optional",
132+ "shape": [
133+ -2
134+ ]
135+ },
136+ {
137+ "name": "scale",
138+ "index": 3,
139+ "dtype": "int64",
140+ "format": "ND",
141+ "paramType": "optional",
142+ "shape": [
143+ -2
144+ ]
145+ }
146+ ],
147+ "outputs": [
148+ {
149+ "name": "y",
150+ "index": 0,
151+ "dtype": "float16",
152+ "format": "ND",
153+ "paramType": "required",
154+ "shape": [
155+ -2
156+ ]
157+ }
158+ ],
159+ "attrs": [
160+ {
161+ "name": "perm_x1",
162+ "dtype": "list_int",
163+ "value": [
164+ 1,
165+ 0,
166+ 2
167+ ]
168+ },
169+ {
170+ "name": "perm_x2",
171+ "dtype": "list_int",
172+ "value": [
173+ 0,
174+ 1,
175+ 2
176+ ]
177+ },
178+ {
179+ "name": "perm_y",
180+ "dtype": "list_int",
181+ "value": [
182+ 1,
183+ 0,
184+ 2
185+ ]
186+ },
187+ {
188+ "name": "enable_hf32",
189+ "dtype": "bool",
190+ "value": false
191+ },
192+ {
193+ "name": "batch_split_factor",
194+ "dtype": "int",
195+ "value": 1
196+ }
197+ ]
198+ },
5 {199 {
6 "bin_filename": "TransposeBatchMatMul_ND_NZ_ND_ND_ND_BF16_BF16_FP32_INT64_BF16",200 "bin_filename": "TransposeBatchMatMul_ND_NZ_ND_ND_ND_BF16_BF16_FP32_INT64_BF16",
7- "simplified_key": "diy,2/29/2/2/2/27/27/27/9/27",201+ "simplified_key": "diy,2/29/2/2/2/27/27/0/9/27",
8 "inputs": [202 "inputs": [
9 {203 {
10 "name": "x1",204 "name": "x1",
@@ -99,6 +293,103 @@
99 }293 }
100 ]294 ]
101 },295 },
296+ {
297+ "bin_filename": "TransposeBatchMatMul_ND_NZ_ND_ND_ND_BF16_BF16_BF16_INT64_BF16",
298+ "simplified_key": "diy,2/29/2/2/2/27/27/27/9/27",
299+ "inputs": [
300+ {
301+ "name": "x1",
302+ "index": 0,
303+ "dtype": "bfloat16",
304+ "format": "ND",
305+ "paramType": "required",
306+ "shape": [
307+ -2
308+ ]
309+ },
310+ {
311+ "name": "x2",
312+ "index": 1,
313+ "dtype": "bfloat16",
314+ "format": "FRACTAL_NZ",
315+ "paramType": "required",
316+ "shape": [
317+ -2
318+ ]
319+ },
320+ {
321+ "name": "bias",
322+ "index": 2,
323+ "dtype": "bfloat16",
324+ "format": "ND",
325+ "paramType": "optional",
326+ "shape": [
327+ -2
328+ ]
329+ },
330+ {
331+ "name": "scale",
332+ "index": 3,
333+ "dtype": "int64",
334+ "format": "ND",
335+ "paramType": "optional",
336+ "shape": [
337+ -2
338+ ]
339+ }
340+ ],
341+ "outputs": [
342+ {
343+ "name": "y",
344+ "index": 0,
345+ "dtype": "bfloat16",
346+ "format": "ND",
347+ "paramType": "required",
348+ "shape": [
349+ -2
350+ ]
351+ }
352+ ],
353+ "attrs": [
354+ {
355+ "name": "perm_x1",
356+ "dtype": "list_int",
357+ "value": [
358+ 1,
359+ 0,
360+ 2
361+ ]
362+ },
363+ {
364+ "name": "perm_x2",
365+ "dtype": "list_int",
366+ "value": [
367+ 0,
368+ 1,
369+ 2
370+ ]
371+ },
372+ {
373+ "name": "perm_y",
374+ "dtype": "list_int",
375+ "value": [
376+ 1,
377+ 0,
378+ 2
379+ ]
380+ },
381+ {
382+ "name": "enable_hf32",
383+ "dtype": "bool",
384+ "value": false
385+ },
386+ {
387+ "name": "batch_split_factor",
388+ "dtype": "int",
389+ "value": 1
390+ }
391+ ]
392+ },
102 {393 {
103 "bin_filename": "TransposeBatchMatMul_ND_ND_ND_ND_ND_FP16_FP16_FP16_INT64_FP16",394 "bin_filename": "TransposeBatchMatMul_ND_ND_ND_ND_ND_FP16_FP16_FP16_INT64_FP16",
104 "simplified_key": "diy,2/2/2/2/2/1/1/1/9/1",395 "simplified_key": "diy,2/2/2/2/2/1/1/1/9/1",
@@ -2,9 +2,203 @@
2 "op_type": "TransposeBatchMatMul",2 "op_type": "TransposeBatchMatMul",
3 "optional_input_mode": "gen_placeholder",3 "optional_input_mode": "gen_placeholder",
4 "op_list": [4 "op_list": [
5+ {
6+ "bin_filename": "TransposeBatchMatMul_ND_NZ_ND_ND_ND_FP16_FP16_FP32_INT64_FP16",
7+ "simplified_key": "diy,2/29/2/2/2/1/1/0/9/1",
8+ "inputs": [
9+ {
10+ "name": "x1",
11+ "index": 0,
12+ "dtype": "float16",
13+ "format": "ND",
14+ "paramType": "required",
15+ "shape": [
16+ -2
17+ ]
18+ },
19+ {
20+ "name": "x2",
21+ "index": 1,
22+ "dtype": "float16",
23+ "format": "FRACTAL_NZ",
24+ "paramType": "required",
25+ "shape": [
26+ -2
27+ ]
28+ },
29+ {
30+ "name": "bias",
31+ "index": 2,
32+ "dtype": "float32",
33+ "format": "ND",
34+ "paramType": "optional",
35+ "shape": [
36+ -2
37+ ]
38+ },
39+ {
40+ "name": "scale",
41+ "index": 3,
42+ "dtype": "int64",
43+ "format": "ND",
44+ "paramType": "optional",
45+ "shape": [
46+ -2
47+ ]
48+ }
49+ ],
50+ "outputs": [
51+ {
52+ "name": "y",
53+ "index": 0,
54+ "dtype": "float16",
55+ "format": "ND",
56+ "paramType": "required",
57+ "shape": [
58+ -2
59+ ]
60+ }
61+ ],
62+ "attrs": [
63+ {
64+ "name": "perm_x1",
65+ "dtype": "list_int",
66+ "value": [
67+ 1,
68+ 0,
69+ 2
70+ ]
71+ },
72+ {
73+ "name": "perm_x2",
74+ "dtype": "list_int",
75+ "value": [
76+ 0,
77+ 1,
78+ 2
79+ ]
80+ },
81+ {
82+ "name": "perm_y",
83+ "dtype": "list_int",
84+ "value": [
85+ 1,
86+ 0,
87+ 2
88+ ]
89+ },
90+ {
91+ "name": "enable_hf32",
92+ "dtype": "bool",
93+ "value": false
94+ },
95+ {
96+ "name": "batch_split_factor",
97+ "dtype": "int",
98+ "value": 1
99+ }
100+ ]
101+ },
102+ {
103+ "bin_filename": "TransposeBatchMatMul_ND_NZ_ND_ND_ND_FP16_FP16_FP16_INT64_FP16",
104+ "simplified_key": "diy,2/29/2/2/2/1/1/1/9/1",
105+ "inputs": [
106+ {
107+ "name": "x1",
108+ "index": 0,
109+ "dtype": "float16",
110+ "format": "ND",
111+ "paramType": "required",
112+ "shape": [
113+ -2
114+ ]
115+ },
116+ {
117+ "name": "x2",
118+ "index": 1,
119+ "dtype": "float16",
120+ "format": "FRACTAL_NZ",
121+ "paramType": "required",
122+ "shape": [
123+ -2
124+ ]
125+ },
126+ {
127+ "name": "bias",
128+ "index": 2,
129+ "dtype": "float16",
130+ "format": "ND",
131+ "paramType": "optional",
132+ "shape": [
133+ -2
134+ ]
135+ },
136+ {
137+ "name": "scale",
138+ "index": 3,
139+ "dtype": "int64",
140+ "format": "ND",
141+ "paramType": "optional",
142+ "shape": [
143+ -2
144+ ]
145+ }
146+ ],
147+ "outputs": [
148+ {
149+ "name": "y",
150+ "index": 0,
151+ "dtype": "float16",
152+ "format": "ND",
153+ "paramType": "required",
154+ "shape": [
155+ -2
156+ ]
157+ }
158+ ],
159+ "attrs": [
160+ {
161+ "name": "perm_x1",
162+ "dtype": "list_int",
163+ "value": [
164+ 1,
165+ 0,
166+ 2
167+ ]
168+ },
169+ {
170+ "name": "perm_x2",
171+ "dtype": "list_int",
172+ "value": [
173+ 0,
174+ 1,
175+ 2
176+ ]
177+ },
178+ {
179+ "name": "perm_y",
180+ "dtype": "list_int",
181+ "value": [
182+ 1,
183+ 0,
184+ 2
185+ ]
186+ },
187+ {
188+ "name": "enable_hf32",
189+ "dtype": "bool",
190+ "value": false
191+ },
192+ {
193+ "name": "batch_split_factor",
194+ "dtype": "int",
195+ "value": 1
196+ }
197+ ]
198+ },
5 {199 {
6 "bin_filename": "TransposeBatchMatMul_ND_NZ_ND_ND_ND_BF16_BF16_FP32_INT64_BF16",200 "bin_filename": "TransposeBatchMatMul_ND_NZ_ND_ND_ND_BF16_BF16_FP32_INT64_BF16",
7- "simplified_key": "diy,2/29/2/2/2/27/27/27/9/27",201+ "simplified_key": "diy,2/29/2/2/2/27/27/0/9/27",
8 "inputs": [202 "inputs": [
9 {203 {
10 "name": "x1",204 "name": "x1",
@@ -99,6 +293,103 @@
99 }293 }
100 ]294 ]
101 },295 },
296+ {
297+ "bin_filename": "TransposeBatchMatMul_ND_NZ_ND_ND_ND_BF16_BF16_BF16_INT64_BF16",
298+ "simplified_key": "diy,2/29/2/2/2/27/27/27/9/27",
299+ "inputs": [
300+ {
301+ "name": "x1",
302+ "index": 0,
303+ "dtype": "bfloat16",
304+ "format": "ND",
305+ "paramType": "required",
306+ "shape": [
307+ -2
308+ ]
309+ },
310+ {
311+ "name": "x2",
312+ "index": 1,
313+ "dtype": "bfloat16",
314+ "format": "FRACTAL_NZ",
315+ "paramType": "required",
316+ "shape": [
317+ -2
318+ ]
319+ },
320+ {
321+ "name": "bias",
322+ "index": 2,
323+ "dtype": "bfloat16",
324+ "format": "ND",
325+ "paramType": "optional",
326+ "shape": [
327+ -2
328+ ]
329+ },
330+ {
331+ "name": "scale",
332+ "index": 3,
333+ "dtype": "int64",
334+ "format": "ND",
335+ "paramType": "optional",
336+ "shape": [
337+ -2
338+ ]
339+ }
340+ ],
341+ "outputs": [
342+ {
343+ "name": "y",
344+ "index": 0,
345+ "dtype": "bfloat16",
346+ "format": "ND",
347+ "paramType": "required",
348+ "shape": [
349+ -2
350+ ]
351+ }
352+ ],
353+ "attrs": [
354+ {
355+ "name": "perm_x1",
356+ "dtype": "list_int",
357+ "value": [
358+ 1,
359+ 0,
360+ 2
361+ ]
362+ },
363+ {
364+ "name": "perm_x2",
365+ "dtype": "list_int",
366+ "value": [
367+ 0,
368+ 1,
369+ 2
370+ ]
371+ },
372+ {
373+ "name": "perm_y",
374+ "dtype": "list_int",
375+ "value": [
376+ 1,
377+ 0,
378+ 2
379+ ]
380+ },
381+ {
382+ "name": "enable_hf32",
383+ "dtype": "bool",
384+ "value": false
385+ },
386+ {
387+ "name": "batch_split_factor",
388+ "dtype": "int",
389+ "value": 1
390+ }
391+ ]
392+ },
102 {393 {
103 "bin_filename": "TransposeBatchMatMul_ND_ND_ND_ND_ND_FP16_FP16_FP16_INT64_FP16",394 "bin_filename": "TransposeBatchMatMul_ND_ND_ND_ND_ND_FP16_FP16_FP16_INT64_FP16",
104 "simplified_key": "diy,2/2/2/2/2/1/1/1/9/1",395 "simplified_key": "diy,2/2/2/2/2/1/1/1/9/1",
@@ -21,29 +21,29 @@ public:
21 {21 {
22 this->Input("x1")22 this->Input("x1")
23 .ParamType(REQUIRED)23 .ParamType(REQUIRED)
24- .DataType({ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_BF16, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_BF16})24+ .DataType({ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_BF16, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_FLOAT16})
25- .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})25+ .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
26- .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});26+ .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
27 this->Input("x2")27 this->Input("x2")
28 .ParamType(REQUIRED)28 .ParamType(REQUIRED)
29- .DataType({ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_BF16, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_BF16})29+ .DataType({ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_BF16, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_FLOAT16})
30- .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_FRACTAL_NZ})30+ .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ})
31- .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_FRACTAL_NZ});31+ .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ});
32 this->Input("bias")32 this->Input("bias")
33 .ParamType(OPTIONAL)33 .ParamType(OPTIONAL)
34- .DataType({ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT})34+ .DataType({ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_BF16, ge::DT_FLOAT, ge::DT_FLOAT16})
35- .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})35+ .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
36- .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});36+ .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
37 this->Input("scale")37 this->Input("scale")
38 .ParamType(OPTIONAL)38 .ParamType(OPTIONAL)
39- .DataType({ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_UINT64, ge::DT_INT64})39+ .DataType({ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_UINT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64})
40- .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})40+ .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
41- .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});41+ .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
42 this->Output("y")42 this->Output("y")
43 .ParamType(REQUIRED)43 .ParamType(REQUIRED)
44- .DataType({ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_BF16, ge::DT_BF16, ge::DT_INT8, ge::DT_INT8, ge::DT_BF16})44+ .DataType({ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_BF16, ge::DT_BF16, ge::DT_INT8, ge::DT_INT8, ge::DT_BF16, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_FLOAT16})
45- .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})45+ .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
46- .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});46+ .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
47 this->Attr("perm_x1")47 this->Attr("perm_x1")
48 .AttrType(OPTIONAL)48 .AttrType(OPTIONAL)
49 .ListInt();49 .ListInt();