已合并
[A2/A3]新增TransposeBatchMatMul算子NZ格式aclnn接口 #1702
fbbccc创建于 2月9日
[A2/A3]新增TransposeBatchMatMul算子NZ格式aclnn接口 #1702
已合并
共 5 个文件变更+796-19
| @@ -36,11 +36,13 @@ using namespace op; | |||
| 36 | using namespace Ops::NN; | 36 | using namespace Ops::NN; |
| 37 | static const std::initializer_list<op::DataType> x1_SUPPORT_LIST = {DataType::DT_FLOAT, DataType::DT_FLOAT16, DataType::DT_BF16}; | 37 | static const std::initializer_list<op::DataType> x1_SUPPORT_LIST = {DataType::DT_FLOAT, DataType::DT_FLOAT16, DataType::DT_BF16}; |
| 38 | static const std::initializer_list<op::DataType> x2_SUPPORT_LIST = {DataType::DT_FLOAT, DataType::DT_FLOAT16, DataType::DT_BF16}; | 38 | static 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}; | ||
| 39 | static const std::initializer_list<op::DataType> x1_SCALE_SUPPORT_LIST = {DataType::DT_FLOAT16}; | 40 | static const std::initializer_list<op::DataType> x1_SCALE_SUPPORT_LIST = {DataType::DT_FLOAT16}; |
| 40 | static const std::initializer_list<op::DataType> x2_SCALE_SUPPORT_LIST = {DataType::DT_FLOAT16}; | 41 | static const std::initializer_list<op::DataType> x2_SCALE_SUPPORT_LIST = {DataType::DT_FLOAT16}; |
| 41 | static const std::initializer_list<op::DataType> SCALE_DTYPE_SUPPORT_LIST = {DataType::DT_INT64, DataType::DT_UINT64}; | 42 | static const std::initializer_list<op::DataType> SCALE_DTYPE_SUPPORT_LIST = {DataType::DT_INT64, DataType::DT_UINT64}; |
| 42 | static const std::initializer_list<op::DataType> OUT_DTYPE_SUPPORT_LIST = {DataType::DT_FLOAT, DataType::DT_FLOAT16, DataType::DT_BF16, DataType::DT_INT8}; | 43 | static const std::initializer_list<op::DataType> OUT_DTYPE_SUPPORT_LIST = {DataType::DT_FLOAT, DataType::DT_FLOAT16, DataType::DT_BF16, DataType::DT_INT8}; |
| 43 | static constexpr size_t EXPECTED_DIM = 3; | 44 | static constexpr size_t EXPECTED_DIM = 3; |
| 45 | +static constexpr int EXPECTED_NZ_DIM = 5; | ||
| 44 | static constexpr int BLOCK_SIZE = 16; | 46 | static constexpr int BLOCK_SIZE = 16; |
| 45 | static constexpr int SUPPORTED_INNER_AXIS = 65536; | 47 | static constexpr int SUPPORTED_INNER_AXIS = 65536; |
| 46 | 48 | ||
| @@ -59,6 +61,43 @@ inline static bool CheckNotNull(const aclTensor* x1, const aclTensor* x2, const | |||
| 59 | inline static bool CheckDtypeValid(const aclTensor* x1, const aclTensor* x2, | 61 | inline static bool CheckDtypeValid(const aclTensor* x1, const aclTensor* x2, |
| 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 | + | ||
| 175 | inline static aclnnStatus CheckParams(const aclTensor* x1, const aclTensor* x2, const aclTensor* scale, aclTensor* out, | 225 | inline 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 | + | ||
| 251 | aclnnStatus aclnnTransposeBatchMatMulGetWorkspaceSize(const aclTensor* x1, const aclTensor* x2, const aclTensor* bias, | 345 | aclnnStatus 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 | ||
这段校验是否可以放到checkParams中,或者封装一个checkParamsWeightNz复用checkParams ![]() ![]() | |||
| 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* | |||
| 49 | ACLNN_API aclnnStatus aclnnTransposeBatchMatMul(void* workspace, uint64_t workspaceSize, aclOpExecutor* executor, | 49 | ACLNN_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 | 88 | ||
| 53 | } | 89 | } |
| 54 | 90 | ||
Mmatmul/transpose_batch_mat_mul/op_host/config/ascend910_93/transpose_batch_mat_mul_binary.json+292-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", |
| @@ -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(); |


在这个分支增加分支