已合并
fix: wqbmmv2 uint8承载antiquantScale修正适配float4_e2m1 weight场景 #10133
fix: wqbmmv2 uint8承载antiquantScale修正适配float4_e2m1 weight场景 #10133
已合并
马琦钧创建于 9月9日
共 2 个文件变更+74-1
@@ -1587,8 +1587,10 @@ static aclnnStatus TensorPreProcess(TupleTensor mandatoryTensors, TupleTensor op
1587 }1587 }
1588 1588 
1589 // microscaling场景,采用uint8承载float8_e8m0数据,此处需修正antiquantScale dtype1589 // 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)
3082TEST_F(l2_weight_quant_batch_matmul_v2_test_950, ascend950_fp8MxMode)3153TEST_F(l2_weight_quant_batch_matmul_v2_test_950, ascend950_fp8MxMode)
3083{3154{