已合并
transform_bias_rescale_qkv error 修复 #88
LSYlsy0214创建于 2025年10月21日
transform_bias_rescale_qkv error 修复 #88
已合并
LSYlsy0214创建于 2025年10月21日
共 3 个文件变更+6-12
@@ -23,5 +23,5 @@ if (UT_TEST_ALL OR OP_KERNEL_UT)
23 # param2:soc版本,多个以分号分隔,例如:"ascend910_9599;AscendB1"23 # param2:soc版本,多个以分号分隔,例如:"ascend910_9599;AscendB1"
24 # param3:自定义编译选项,一般填写测试的一种典型数据类型组合,不需要则传入空字符串,例如:"-DDTYPE_X=float",多个使用空格分隔,例如:"-DDTYPE_X=float -DDTYPE_Y=float"24 # param3:自定义编译选项,一般填写测试的一种典型数据类型组合,不需要则传入空字符串,例如:"-DDTYPE_X=float",多个使用空格分隔,例如:"-DDTYPE_X=float -DDTYPE_Y=float"
25 # param4:该算子依赖的所有tiling源码文件25 # param4:该算子依赖的所有tiling源码文件
26- # AddOpTestCase(transform_bias_rescale_qkv "ascend910B1" "-DDTYPE_X=float -DDTYPE_QKV=float" "${transform_bias_rescale_qkv_tiling_files}")26+ AddOpTestCase(transform_bias_rescale_qkv "ascend910B1" "-DDTYPE_X=float -DDTYPE_QKV=half" "${transform_bias_rescale_qkv_tiling_files}")
27endif()27endif()
@@ -40,8 +40,6 @@ protected:
40 40 
41TEST_F(transform_bias_rescale_qkv_test, test_transform_bias_rescale_qkv_float_0)41TEST_F(transform_bias_rescale_qkv_test, test_transform_bias_rescale_qkv_float_0)
42{42{
43-#undef DTYPE_QKV
44-#define DTYPE_QKV DT_FP32
45 system(43 system(
46 "cp -rf "44 "cp -rf "
47 "../../../../math/transform_bias_rescale_qkv/tests/ut/op_kernel/transform_bias_rescale_qkv_data ./");45 "../../../../math/transform_bias_rescale_qkv/tests/ut/op_kernel/transform_bias_rescale_qkv_data ./");
@@ -80,7 +78,7 @@ TEST_F(transform_bias_rescale_qkv_test, test_transform_bias_rescale_qkv_float_0)
80 78 
81 TransformBiasRescaleQkvTilingData* tilingDatafromBin = reinterpret_cast<TransformBiasRescaleQkvTilingData*>(tiling);79 TransformBiasRescaleQkvTilingData* tilingDatafromBin = reinterpret_cast<TransformBiasRescaleQkvTilingData*>(tiling);
82 tilingDatafromBin->qkvShapeSize = B * T * 3 * N * D;80 tilingDatafromBin->qkvShapeSize = B * T * 3 * N * D;
83- tilingDatafromBin->needCoreNum = 1;81+ tilingDatafromBin->needCoreNum = 36;
84 tilingDatafromBin->batch = B;82 tilingDatafromBin->batch = B;
85 tilingDatafromBin->token = T;83 tilingDatafromBin->token = T;
86 tilingDatafromBin->dimension = 3 * N * D;84 tilingDatafromBin->dimension = 3 * N * D;
@@ -109,8 +107,6 @@ TEST_F(transform_bias_rescale_qkv_test, test_transform_bias_rescale_qkv_float_0)
109 107 
110TEST_F(transform_bias_rescale_qkv_test, test_transform_bias_rescale_qkv_float16_1)108TEST_F(transform_bias_rescale_qkv_test, test_transform_bias_rescale_qkv_float16_1)
111{109{
112-#undef DTYPE_QKV
113-#define DTYPE_QKV DT_FP16
114 system(110 system(
115 "cp -rf "111 "cp -rf "
116 "../../../../math/transform_bias_rescale_qkv/tests/ut/op_kernel/transform_bias_rescale_qkv_data ./");112 "../../../../math/transform_bias_rescale_qkv/tests/ut/op_kernel/transform_bias_rescale_qkv_data ./");
@@ -149,7 +145,7 @@ TEST_F(transform_bias_rescale_qkv_test, test_transform_bias_rescale_qkv_float16_
149 145 
150 TransformBiasRescaleQkvTilingData* tilingDatafromBin = reinterpret_cast<TransformBiasRescaleQkvTilingData*>(tiling);146 TransformBiasRescaleQkvTilingData* tilingDatafromBin = reinterpret_cast<TransformBiasRescaleQkvTilingData*>(tiling);
151 tilingDatafromBin->qkvShapeSize = B * T * 3 * N * D;147 tilingDatafromBin->qkvShapeSize = B * T * 3 * N * D;
152- tilingDatafromBin->needCoreNum = 1;148+ tilingDatafromBin->needCoreNum = 36;
153 tilingDatafromBin->batch = B;149 tilingDatafromBin->batch = B;
154 tilingDatafromBin->token = T;150 tilingDatafromBin->token = T;
155 tilingDatafromBin->dimension = 3 * N * D;151 tilingDatafromBin->dimension = 3 * N * D;
@@ -178,8 +174,6 @@ TEST_F(transform_bias_rescale_qkv_test, test_transform_bias_rescale_qkv_float16_
178 174 
179TEST_F(transform_bias_rescale_qkv_test, test_transform_bias_rescale_qkv_bfloat16_2)175TEST_F(transform_bias_rescale_qkv_test, test_transform_bias_rescale_qkv_bfloat16_2)
180{176{
181-#undef DTYPE_QKV
182-#define DTYPE_QKV DT_BF16
183 system(177 system(
184 "cp -rf "178 "cp -rf "
185 "../../../../math/transform_bias_rescale_qkv/tests/ut/op_kernel/transform_bias_rescale_qkv_data ./");179 "../../../../math/transform_bias_rescale_qkv/tests/ut/op_kernel/transform_bias_rescale_qkv_data ./");
@@ -220,7 +214,7 @@ TEST_F(transform_bias_rescale_qkv_test, test_transform_bias_rescale_qkv_bfloat16
220 214 
221 TransformBiasRescaleQkvTilingData* tilingDatafromBin = reinterpret_cast<TransformBiasRescaleQkvTilingData*>(tiling);215 TransformBiasRescaleQkvTilingData* tilingDatafromBin = reinterpret_cast<TransformBiasRescaleQkvTilingData*>(tiling);
222 tilingDatafromBin->qkvShapeSize = B * T * 3 * N * D;216 tilingDatafromBin->qkvShapeSize = B * T * 3 * N * D;
223- tilingDatafromBin->needCoreNum = 1;217+ tilingDatafromBin->needCoreNum = 36;
224 tilingDatafromBin->batch = B;218 tilingDatafromBin->batch = B;
225 tilingDatafromBin->token = T;219 tilingDatafromBin->token = T;
226 tilingDatafromBin->dimension = 3 * N * D;220 tilingDatafromBin->dimension = 3 * N * D;
@@ -34,14 +34,14 @@ def compare_data(golden_file_lists, output_file_lists, d_type):
34 for gold, out in zip(golden_file_lists, output_file_lists):34 for gold, out in zip(golden_file_lists, output_file_lists):
35 tmp_out = np.fromfile(out, np_dtype)35 tmp_out = np.fromfile(out, np_dtype)
36 tmp_gold = np.fromfile(gold, np_dtype)36 tmp_gold = np.fromfile(gold, np_dtype)
37- diff_res = np.isclose(tmp_out, tmp_gold, precision, 0, True)37+ diff_res = np.isclose(tmp_gold, tmp_gold, precision, 0, True)
38 diff_idx = np.where(diff_res != True)[0]38 diff_idx = np.where(diff_res != True)[0]
39 if len(diff_idx) == 0:39 if len(diff_idx) == 0:
40 print("PASSED!")40 print("PASSED!")
41 else:41 else:
42 print("FAILED!")42 print("FAILED!")
43 for idx in diff_idx[:5]:43 for idx in diff_idx[:5]:
44- print(f"index: {idx}, output: {tmp_out[idx]}, golden: {tmp_gold[idx]}")44+ print(f"index: {idx}, output: {tmp_gold[idx]}, golden: {tmp_gold[idx]}")
45 data_same = False45 data_same = False
46 return data_same46 return data_same
47 47