已合并
fix: wqbmmv2 uint8承载antiquantScale修正适配float4_e2m1 weight场景 #10133
马琦钧创建于 9月9日
fix: wqbmmv2 uint8承载antiquantScale修正适配float4_e2m1 weight场景 #10133
已合并
共 2 个文件变更+74-1
| @@ -1587,8 +1587,10 @@ static aclnnStatus TensorPreProcess(TupleTensor mandatoryTensors, TupleTensor op | |||
| 1587 | } | 1587 | } |
| 1588 | 1588 | ||
| 1589 | // microscaling场景,采用uint8承载float8_e8m0数据,此处需修正antiquantScale dtype | 1589 | // microscaling场景,采用uint8承载float8_e8m0数据,此处需修正antiquantScale dtype |
| 1590 | + // weight为float32承载float4_e2m1或直接为float4_e2m1时均支持 | ||
| 1590 | if (op::GetCurrentPlatformInfo().GetCurNpuArch() == NpuArch::DAV_3510 && | 1591 | if (op::GetCurrentPlatformInfo().GetCurNpuArch() == NpuArch::DAV_3510 && |
| 1591 | - weight->GetDataType() == DataType::DT_FLOAT && antiquantScaleRef->GetDataType() == DataType::DT_UINT8) { | 1592 | + (weight->GetDataType() == DataType::DT_FLOAT || weight->GetDataType() == DataType::DT_FLOAT4_E2M1) && |
| 1593 | + antiquantScaleRef->GetDataType() == DataType::DT_UINT8) { | ||
| 1592 | CHECK_RET(ModifyTensorDtype(antiquantScaleRef, nullptr, DataType::DT_FLOAT8_E8M0, executor) == ACLNN_SUCCESS, | 1594 | CHECK_RET(ModifyTensorDtype(antiquantScaleRef, nullptr, DataType::DT_FLOAT8_E8M0, executor) == ACLNN_SUCCESS, |
| 1593 | ACLNN_ERR_PARAM_INVALID); | 1595 | ACLNN_ERR_PARAM_INVALID); |
| 1594 | OP_LOGD("The conversion of antiquantScale from uint8 to fp8 is completed."); | 1596 | OP_LOGD("The conversion of antiquantScale from uint8 to fp8 is completed."); |
| @@ -2605,6 +2605,62 @@ static WeightQuantBatchMatmulV2TestParam casesParamsAscend950[] = { | |||
| 2605 | CONTIGUOUS, | 2605 | CONTIGUOUS, |
| 2606 | CONTIGUOUS, | 2606 | CONTIGUOUS, |
| 2607 | CONTIGUOUS}, | 2607 | CONTIGUOUS}, |
| 2608 | + {"Ascend950_case_a16mxf4_nd_weight_fp4_uint8_scale", | ||
| 2609 | + {2, 64}, | ||
| 2610 | + {64, 128}, | ||
| 2611 | + {2, 128}, | ||
| 2612 | + {2, 128}, | ||
| 2613 | + {1, 128}, | ||
| 2614 | + {1, 128}, | ||
| 2615 | + {1, 128}, | ||
| 2616 | + 32, | ||
| 2617 | + {2, 128}, | ||
| 2618 | + ACL_FLOAT16, | ||
| 2619 | + ACL_FLOAT4_E2M1, | ||
| 2620 | + ACL_UINT8, | ||
| 2621 | + ACL_FLOAT16, | ||
| 2622 | + ACL_UINT64, | ||
| 2623 | + ACL_FLOAT, | ||
| 2624 | + ACL_FLOAT16, | ||
| 2625 | + ACL_FLOAT16, | ||
| 2626 | + ACL_FORMAT_ND, | ||
| 2627 | + ACL_FORMAT_ND, | ||
| 2628 | + false, | ||
| 2629 | + false, | ||
| 2630 | + false, | ||
| 2631 | + false, | ||
| 2632 | + ACLNN_SUCCESS, | ||
| 2633 | + CONTIGUOUS, | ||
| 2634 | + CONTIGUOUS, | ||
| 2635 | + CONTIGUOUS}, | ||
| 2636 | + {"Ascend950_case_a16mxf4_nd_weight_fp32_uint8_scale", | ||
| 2637 | + {2, 64}, | ||
| 2638 | + {64, 16}, // weight N=16 is packed (8 FP4 in 1 FP32), host unpacks to logical N=128 | ||
| 2639 | + {2, 128}, | ||
| 2640 | + {2, 128}, | ||
| 2641 | + {1, 128}, | ||
| 2642 | + {1, 128}, | ||
| 2643 | + {1, 128}, | ||
| 2644 | + 32, | ||
| 2645 | + {2, 128}, | ||
| 2646 | + ACL_FLOAT16, | ||
| 2647 | + ACL_FLOAT, | ||
| 2648 | + ACL_UINT8, | ||
| 2649 | + ACL_FLOAT16, | ||
| 2650 | + ACL_UINT64, | ||
| 2651 | + ACL_FLOAT, | ||
| 2652 | + ACL_FLOAT16, | ||
| 2653 | + ACL_FLOAT16, | ||
| 2654 | + ACL_FORMAT_ND, | ||
| 2655 | + ACL_FORMAT_ND, | ||
| 2656 | + false, | ||
| 2657 | + false, | ||
| 2658 | + false, | ||
| 2659 | + false, | ||
| 2660 | + ACLNN_SUCCESS, | ||
| 2661 | + CONTIGUOUS, | ||
| 2662 | + CONTIGUOUS, | ||
| 2663 | + CONTIGUOUS}, | ||
| 2608 | {"Ascend950_case_a16mxf4_nd_invalid_group_size", | 2664 | {"Ascend950_case_a16mxf4_nd_invalid_group_size", |
| 2609 | {2, 64}, | 2665 | {2, 64}, |
| 2610 | {64, 128}, | 2666 | {64, 128}, |
| @@ -3078,6 +3134,21 @@ TEST_F(l2_weight_quant_batch_matmul_v2_test_950, ascend950_fp4PerChannel) | |||
| 3078 | EXPECT_NE(ret, ACLNN_SUCCESS); | 3134 | EXPECT_NE(ret, ACLNN_SUCCESS); |
| 3079 | } | 3135 | } |
| 3080 | 3136 | ||
| 3137 | +// 950: MX A16F4 ND——uint8 承载 E8M0 antiquantScale 的 dtype 修正在 V3 入口同样生效 | ||
| 3138 | +// (与参数化用例 Ascend950_case_a16mxf4_nd_weight_fp4_uint8_scale 同组合,走 V3 入口) | ||
| 3139 | +TEST_F(l2_weight_quant_batch_matmul_v2_test_950, ascend950_a16mxf4Fp4Uint8ScaleV3) | ||
| 3140 | +{ | ||
| 3141 | + auto x = CreateTensorDesc({2, 64}, ACL_FLOAT16, ACL_FORMAT_ND, CONTIGUOUS); | ||
| 3142 | + auto weight = CreateTensorDesc({64, 128}, ACL_FLOAT4_E2M1, ACL_FORMAT_ND, CONTIGUOUS); | ||
| 3143 | + auto scale = CreateTensorDesc({2, 128}, ACL_UINT8, ACL_FORMAT_ND, CONTIGUOUS); | ||
| 3144 | + auto y = CreateTensorDesc({2, 128}, ACL_FLOAT16, ACL_FORMAT_ND, CONTIGUOUS); | ||
| 3145 | + uint64_t ws = 0; | ||
| 3146 | + aclOpExecutor* exe = nullptr; | ||
| 3147 | + auto ret = aclnnWeightQuantBatchMatmulV3GetWorkspaceSize(x, weight, scale, nullptr, nullptr, nullptr, nullptr, 32, | ||
| 3148 | + 0, y, &ws, &exe); | ||
| 3149 | + EXPECT_EQ(ret, ACLNN_SUCCESS); | ||
| 3150 | +} | ||
| 3151 | + | ||
| 3081 | // 950: FP8 weight with MX mode (antiquantScale=FLOAT8_E8M0) | 3152 | // 950: FP8 weight with MX mode (antiquantScale=FLOAT8_E8M0) |
| 3082 | TEST_F(l2_weight_quant_batch_matmul_v2_test_950, ascend950_fp8MxMode) | 3153 | TEST_F(l2_weight_quant_batch_matmul_v2_test_950, ascend950_fp8MxMode) |
| 3083 | { | 3154 | { |