已合并
fix: WeightQuantPreprocess A16MXF4 支持转置 weight ND 直拷 #5803
fix: WeightQuantPreprocess A16MXF4 支持转置 weight ND 直拷 #5803
已合并
马琦钧创建于 8月31日
共 4 个文件变更+195-122
@@ -10,9 +10,11 @@
10 10 
11该接口针对Matmul类算子的伪量化参数进行预处理,目前支持处理的数据流如下:11该接口针对Matmul类算子的伪量化参数进行预处理,目前支持处理的数据流如下:
12 12 
13-- MM_MX_A8W4数据流:表示QuantBatchMatmulV4算子的MX量化场景,A8W4表示左矩阵为`float8_e4m3`,右矩阵为`float4_e2m1`。13+- MM_MX_A8W4数据流:表示QuantBatchMatmulV4算子的MX量化场景,A8W4表示左矩阵为`torch.float8_e4m3fn`,右矩阵为`torch_npu.float4_e2m1fn_x2`。
14-- GMM_MX_A8W4数据流:表示GroupedMatmul算子的MX量化场景,A8W4表示左矩阵为`float8_e4m3`,右矩阵为`float4_e2m1`。14+- GMM_MX_A8W4数据流:表示GroupedMatmul算子的MX量化场景,A8W4表示左矩阵为`torch.float8_e4m3fn`,右矩阵为`torch_npu.float4_e2m1fn_x2`。
15-- MM_A16S4数据流:表示WeightQuantBatchMatmulV2算子的pertensor/perchannel/pergroup量化场景,A16S4表示左矩阵为`float16`/`bfloat16`,右矩阵为`int4`。15+- MM_A16S4数据流:表示WeightQuantBatchMatmulV2算子的pertensor/perchannel/pergroup量化场景,A16S4表示左矩阵为`torch.float16`/`torch.bfloat16`,右矩阵为`torch_npu.int4`。
16+- MM_A16F4数据流:表示WeightQuantBatchMatmulV2算子的pergroup量化场景,A16F4表示左矩阵为`torch.float16`/`torch.bfloat16`,右矩阵为`torch_npu.float4_e2m1fn_x2`(使用`torch.uint8`承载)。
17+- MM_MX_A16F4数据流:表示WeightQuantBatchMatmulV2算子的MX量化场景,A16F4表示左矩阵为`torch.float16`/`torch.bfloat16`,右矩阵为`torch_npu.float4_e2m1fn_x2`(使用`torch.uint8`承载)。
16 18 
17## 函数原型19## 函数原型
18 20 
@@ -23,27 +25,27 @@ torch_npu.npu_weight_quant_preprocess(weight, weight_scale, x_dtype, weight_dtyp
23## 参数说明25## 参数说明
24 26 
25- **weight**(`Tensor`):**必选参数**,Matmul的权重矩阵,支持非连续`Tensor`。27- **weight**(`Tensor`):**必选参数**,Matmul的权重矩阵,支持非连续`Tensor`。
26- - 逻辑数据类型支持`int4`(使用`uint8`承载)、`float4_e2m1fn_x2`(使用`uint8`承载),数据格式支持$ND$。支持2维或3维输入,逻辑shape分别为$(K, N)$、$(G, K, N)$,其中$G$表示GroupedMatmul算子的G轴。1个`uint8`元素打包2个4-bit数据,要求打包维度的元素个数为偶数:沿K维打包时`weight`的物理K维长度为$K/2$,沿N维打包时物理N维长度为$N/2$。28+ - 逻辑数据类型支持`torch_npu.int4`(使用`torch.uint8`承载)、`torch_npu.float4_e2m1fn_x2`(使用`torch.uint8`承载),数据格式支持$ND$。支持2维或3维输入,逻辑shape分别为$(K, N)$、$(G, K, N)$,其中$G$表示GroupedMatmul算子的G轴。1个`torch.uint8`元素打包2个4-bit数据,要求打包维度的元素个数为偶数:沿K维打包时`weight`的物理K维长度为$K/2$,沿N维打包时物理N维长度为$N/2$。
27- **weight\_scale**(`Tensor`):**必选参数**,权重的反量化scale参数,支持非连续`Tensor`。29- **weight\_scale**(`Tensor`):**必选参数**,权重的反量化scale参数,支持非连续`Tensor`。
28- - 逻辑数据类型支持`float8_e8m0fnu`(使用`uint8`承载)、`float16`、`bfloat16`,数据格式支持$ND$、$NCL$。30+ - 逻辑数据类型支持`torch.float8_e8m0fnu`(使用`torch.uint8`承载)、`torch.float16`、`torch.bfloat16`,数据格式支持$ND$、$NCL$。
29- - MX量化场景:支持3维或4维输入,shape分别为$(ceil\_div(K, 64), N, 2)$、$(G, ceil\_div(K, 64), N, 2)$,其中$G$表示GroupedMatmul算子的G轴。31+ - MX量化场景:A8W4数据流支持3维或4维输入,shape分别为$(ceil\_div(K, 64), N, 2)$、$(G, ceil\_div(K, 64), N, 2)$,其中$G$表示GroupedMatmul算子的G轴;A16F4数据流支持2维输入,shape为$(ceil\_div(K, 32), N)$。
30 - pertensor量化场景:shape为$(1)$、$(1, 1)$。32 - pertensor量化场景:shape为$(1)$、$(1, 1)$。
31 - perchannel量化场景:shape为$(N)$、$(1, N)$。33 - perchannel量化场景:shape为$(N)$、$(1, N)$。
32 - pergroup量化场景:shape为$(G, N)$,其中$G$表示group数量。34 - pergroup量化场景:shape为$(G, N)$,其中$G$表示group数量。
33- **x\_dtype**(`int`):**必选参数**,Matmul的激活矩阵的数据类型。35- **x\_dtype**(`int`):**必选参数**,Matmul的激活矩阵的数据类型。
34- - A16S4数据流支持`torch.float16`、`torch.bfloat16`;A8W4数据流支持`torch.float8_e4m3fn`。36+ - A16S4/A16F4数据流支持`torch.float16`、`torch.bfloat16`;A8W4数据流支持`torch.float8_e4m3fn`。
35- **weight\_dtype**(`int`):**必选参数**,用于指定`weight`中实际承载的数据类型。37- **weight\_dtype**(`int`):**必选参数**,用于指定`weight`中实际承载的数据类型。
36- - A16S4数据流取值为`torch_npu.int4`;其余数据流取值为`torch.float4_e2m1fn_x2`。38+ - A16S4数据流取值为`torch_npu.int4`;其余数据流取值为`torch_npu.float4_e2m1fn_x2`。
37- **weight\_scale\_dtype**(`int`):**必选参数**,用于指定`weight_scale`中实际承载的数据类型,取值需与`weight_scale`的实际数据类型一致。39- **weight\_scale\_dtype**(`int`):**必选参数**,用于指定`weight_scale`中实际承载的数据类型,取值需与`weight_scale`的实际数据类型一致。
38 - 支持`torch.float8_e8m0fnu`、`torch.float16`、`torch.bfloat16`。40 - 支持`torch.float8_e8m0fnu`、`torch.float16`、`torch.bfloat16`。
39- **weight\_offset**(`Tensor`):**可选参数**,权重的反量化offset参数,默认值为`None`。41- **weight\_offset**(`Tensor`):**可选参数**,权重的反量化offset参数,默认值为`None`。
40 - A16S4各数据流支持透传,其shape和数据类型需与`weight_scale`保持一致;其余数据流请传入`None`。42 - A16S4各数据流支持透传,其shape和数据类型需与`weight_scale`保持一致;其余数据流请传入`None`。
41- **bias**(`Tensor`):**可选参数**,Matmul的偏置矩阵,必须为连续的`Tensor`。43- **bias**(`Tensor`):**可选参数**,Matmul的偏置矩阵,必须为连续的`Tensor`。
42- - 数据类型支持`float16`、`bfloat16`、`float32`,数据格式支持$ND$。支持1维或2维输入,shape为$(N)$、$(1, N)$或$(G, N)$。44+ - 数据类型支持`torch.float16`、`torch.bfloat16`、`torch.float32`,数据格式支持$ND$。支持1维或2维输入,shape为$(N)$、$(1, N)$或$(G, N)$。
43- **x\_scale\_dtype**(`int`):**可选参数**,激活的量化scale参数的数据类型,默认值为`None`。45- **x\_scale\_dtype**(`int`):**可选参数**,激活的量化scale参数的数据类型,默认值为`None`。
44- - A8W4数据流仅支持`torch.float8_e8m0fnu`;A16S4数据流请传入`None`。46+ - A8W4数据流仅支持`torch.float8_e8m0fnu`;A16S4/A16F4数据流请传入`None`。
45- **k\_group\_size**(`int`):**可选参数**,权重在pergroup量化时K维度的group大小,默认值为`0`。47- **k\_group\_size**(`int`):**可选参数**,权重在pergroup量化时K维度的group大小,默认值为`0`。
46- - A8W4数据流MX量化场景下取值为`32`;A16S4 pergroup量化场景取值为大于`0`的整数;其余场景使用默认值`0`。48+ - A8W4/A16F4数据流MX量化场景下取值为`32`;A16S4/A16F4数据流pergroup量化场景取值为大于`0`的整数;其余场景使用默认值`0`。
47 49 
48## 返回值说明50## 返回值说明
49 51 
@@ -55,6 +57,8 @@ torch_npu.npu_weight_quant_preprocess(weight, weight_scale, x_dtype, weight_dtyp
55 - A16S4数据流下pertensor(`weight`转置/非转置)、perchannel(`weight`转置)、pergroup(`weight`转置)量化场景:`weight`的数据格式仍为$ND$。57 - A16S4数据流下pertensor(`weight`转置/非转置)、perchannel(`weight`转置)、pergroup(`weight`转置)量化场景:`weight`的数据格式仍为$ND$。
56 - A16S4数据流下perchannel(`weight`非转置)、pergroup(`weight`非转置)量化场景:`weight`的数据格式由$ND$转为$FRACTAL\_NZ\_C0\_8$。58 - A16S4数据流下perchannel(`weight`非转置)、pergroup(`weight`非转置)量化场景:`weight`的数据格式由$ND$转为$FRACTAL\_NZ\_C0\_8$。
57 - A8W4数据流下MX量化场景:`weight`的数据格式由$ND$转为$FRACTAL\_NZ\_C0\_16$。59 - A8W4数据流下MX量化场景:`weight`的数据格式由$ND$转为$FRACTAL\_NZ\_C0\_16$。
60+ - A16F4数据流下pergroup量化场景(仅支持`weight`非转置):`weight`的数据格式由$ND$转为$FRACTAL\_NZ\_C0\_8$。
61+ - A16F4数据流下MX量化场景:`weight`转置输入时输出为$ND$直拷,与输入保持相同的sizes/strides;非转置输入时由$ND$转为$FRACTAL\_NZ\_C0\_8$。
58- **out\_weight\_scale**(`Tensor`):预处理后的`weight_scale`,数据类型与输入`weight_scale`相同。62- **out\_weight\_scale**(`Tensor`):预处理后的`weight_scale`,数据类型与输入`weight_scale`相同。
59- **out\_weight\_offset**(`Tensor`):预处理后的`weight_offset`,未传入`weight_offset`时返回空`Tensor`;A16S4数据流传入时透传,与输入保持相同的sizes/strides。63- **out\_weight\_offset**(`Tensor`):预处理后的`weight_offset`,未传入`weight_offset`时返回空`Tensor`;A16S4数据流传入时透传,与输入保持相同的sizes/strides。
60- **out\_bias**(`Tensor`):预处理后的`bias`,数据类型与输入`bias`相同(若提供了`bias`)。64- **out\_bias**(`Tensor`):预处理后的`bias`,数据类型与输入`bias`相同(若提供了`bias`)。
@@ -64,34 +68,50 @@ torch_npu.npu_weight_quant_preprocess(weight, weight_scale, x_dtype, weight_dtyp
64当前支持如下参数组合:68当前支持如下参数组合:
65 69 
66- **MM_MX_A8W4数据流**:70- **MM_MX_A8W4数据流**:
67- - `weight`数据类型必须为`float4_e2m1`,数据格式为$ND$,K必须满足$K \% k\_group\_size = 0$。71+ - `weight`数据类型必须为`torch_npu.float4_e2m1fn_x2`,数据格式为$ND$,K必须满足$K \% k\_group\_size = 0$。
68- - `weight`的逻辑shape为$\{K, N\}$;使用`uint8`承载fp4数据时,view shape为$\{K/2, N\}$,storage shape为$\{N, K/2\}$(transposed),stride为$[1, K/2]$。72+ - `weight`的逻辑shape为$\{K, N\}$;使用`torch.uint8`承载fp4数据时,view shape为$\{K/2, N\}$,storage shape为$\{N, K/2\}$(transposed),stride为$[1, K/2]$。
69- - `weight_scale`数据类型必须为`float8_e8m0`,数据格式为$ND$/$NCL$。73+ - `weight_scale`数据类型必须为`torch.float8_e8m0fnu`,数据格式为$ND$/$NCL$。
70 - `weight_scale`的view shape为$\{ceil\_div(K, 64), N, 2\}$,storage shape为$\{N, ceil\_div(K, 64), 2\}$(transposed)。74 - `weight_scale`的view shape为$\{ceil\_div(K, 64), N, 2\}$,storage shape为$\{N, ceil\_div(K, 64), 2\}$(transposed)。
71 - `k_group_size`必须等于`32`。75 - `k_group_size`必须等于`32`。
72- - `x_dtype`必须为`float8_e4m3fn`。76+ - `x_dtype`必须为`torch.float8_e4m3fn`。
73- - `x_scale_dtype`必须为`float8_e8m0`。77+ - `x_scale_dtype`必须为`torch.float8_e8m0fnu`。
74 - 当前不支持`weight_offset`,必须传入`None`。78 - 当前不支持`weight_offset`,必须传入`None`。
75- **GMM_MX_A8W4数据流**:79- **GMM_MX_A8W4数据流**:
76- - `weight`数据类型必须为`float4_e2m1`,数据格式为$ND$,K必须满足$K \% k\_group\_size = 0$。80+ - `weight`数据类型必须为`torch_npu.float4_e2m1fn_x2`,数据格式为$ND$,K必须满足$K \% k\_group\_size = 0$。
77- - `weight`的逻辑shape为$\{G, K, N\}$;使用`uint8`承载fp4数据时,view shape为$\{G, K/2, N\}$,storage shape为$\{G, N, K/2\}$(transposed),stride为$[K * N/2, 1, K/2]$。81+ - `weight`的逻辑shape为$\{G, K, N\}$;使用`torch.uint8`承载fp4数据时,view shape为$\{G, K/2, N\}$,storage shape为$\{G, N, K/2\}$(transposed),stride为$[K * N/2, 1, K/2]$。
78- - `weight_scale`数据类型必须为`float8_e8m0`,数据格式为$ND$/$NCL$。82+ - `weight_scale`数据类型必须为`torch.float8_e8m0fnu`,数据格式为$ND$/$NCL$。
79 - `weight_scale`的view shape为$\{G, ceil\_div(K, 64), N, 2\}$,storage shape为$\{G, N, ceil\_div(K, 64), 2\}$(transposed),stride为$[N * ceil\_div(K, 64) * 2, 2, ceil\_div(K, 64) * 2, 1]$。83 - `weight_scale`的view shape为$\{G, ceil\_div(K, 64), N, 2\}$,storage shape为$\{G, N, ceil\_div(K, 64), 2\}$(transposed),stride为$[N * ceil\_div(K, 64) * 2, 2, ceil\_div(K, 64) * 2, 1]$。
80 - `k_group_size`必须等于`32`。84 - `k_group_size`必须等于`32`。
81- - `x_dtype`必须为`float8_e4m3fn`。85+ - `x_dtype`必须为`torch.float8_e4m3fn`。
82- - `x_scale_dtype`必须为`float8_e8m0`。86+ - `x_scale_dtype`必须为`torch.float8_e8m0fnu`。
83 - 当前不支持`weight_offset`,必须传入`None`。87 - 当前不支持`weight_offset`,必须传入`None`。
84- **MM_A16S4数据流**:88- **MM_A16S4数据流**:
85- - `weight`数据类型必须为`int4`,使用`uint8`承载(1个`uint8`元素打包2个int4数据),数据格式为$ND$,支持2维输入,转置与非转置均支持。89+ - `weight`数据类型必须为`torch_npu.int4`,使用`torch.uint8`承载(1个`torch.uint8`元素打包2个int4数据),数据格式为$ND$,支持2维输入,转置与非转置均支持。
86 - 非转置时`weight`的shape为$\{K, N/2\}$(沿N维打包),转置时shape为$\{K/2, N\}$(沿K维打包),打包维度的元素个数须为偶数。90 - 非转置时`weight`的shape为$\{K, N/2\}$(沿N维打包),转置时shape为$\{K/2, N\}$(沿K维打包),打包维度的元素个数须为偶数。
87- - `weight_scale`数据类型为`float16`或`bfloat16`,其shape决定量化粒度。91+ - `weight_scale`数据类型为`torch.float16`或`torch.bfloat16`,其shape决定量化粒度。
88- - `x_dtype`必须为`float16`或`bfloat16`。92+ - `x_dtype`必须为`torch.float16`或`torch.bfloat16`。
89 - `x_scale_dtype`必须传入`None`。93 - `x_scale_dtype`必须传入`None`。
90 - `weight_offset`支持透传,shape和数据类型需与`weight_scale`保持一致。94 - `weight_offset`支持透传,shape和数据类型需与`weight_scale`保持一致。
91 - 根据`weight_scale`的shape分为如下三种场景:95 - 根据`weight_scale`的shape分为如下三种场景:
92 - pertensor:`weight_scale`仅含单个元素,shape为$\{1\}$或$\{1, 1\}$;`k_group_size`使用默认值`0`;`out_weight`为$ND$直拷输出,与输入`weight`保持相同的sizes/strides。96 - pertensor:`weight_scale`仅含单个元素,shape为$\{1\}$或$\{1, 1\}$;`k_group_size`使用默认值`0`;`out_weight`为$ND$直拷输出,与输入`weight`保持相同的sizes/strides。
93 - perchannel:`weight_scale`的shape为$\{N\}$或$\{1, N\}$;`k_group_size`使用默认值`0`;转置输入时`out_weight`为$ND$直拷输出,非转置输入时`out_weight`为$FRACTAL\_NZ\_C0\_8$输出,storage shape为$\{ceil\_div(N, 16), ceil\_div(K, 16), 16, 8\}$(N块在前)。97 - perchannel:`weight_scale`的shape为$\{N\}$或$\{1, N\}$;`k_group_size`使用默认值`0`;转置输入时`out_weight`为$ND$直拷输出,非转置输入时`out_weight`为$FRACTAL\_NZ\_C0\_8$输出,storage shape为$\{ceil\_div(N, 16), ceil\_div(K, 16), 16, 8\}$(N块在前)。
94 - pergroup:`weight_scale`的shape为$\{G, N\}$,其中$G$为group数量且$G > 1$;`k_group_size`必须大于`0`,表示K维pergroup的group大小;输出路由与perchannel场景一致;`weight`转置输入时,`weight_scale`需与`weight`布局一致(同为转置视图)。98 - pergroup:`weight_scale`的shape为$\{G, N\}$,其中$G$为group数量且$G > 1$;`k_group_size`必须大于`0`,表示K维pergroup的group大小;输出路由与perchannel场景一致;`weight`转置输入时,`weight_scale`需与`weight`布局一致(同为转置视图)。
99+- **MM_A16F4数据流**:
100+ - `weight`数据类型必须为`torch_npu.float4_e2m1fn_x2`,使用`torch.uint8`承载(1个`torch.uint8`元素打包2个fp4数据),数据格式为$ND$,支持2维输入,仅支持非转置,shape为$\{K, N/2\}$(沿N维打包),打包维度的元素个数须为偶数。
101+ - `weight_scale`数据类型为`torch.float16`或`torch.bfloat16`,shape为$\{G, N\}$,其中$G$为group数量且$G > 1$。
102+ - `k_group_size`必须大于`0`,表示K维pergroup的group大小。
103+ - `x_dtype`必须为`torch.float16`或`torch.bfloat16`。
104+ - `x_scale_dtype`必须传入`None`。
105+ - 当前不支持`weight_offset`,必须传入`None`。
106+ - `out_weight`为$FRACTAL\_NZ\_C0\_8$输出,storage shape为$\{ceil\_div(N, 16), ceil\_div(K, 16), 16, 8\}$(N块在前)。
107+- **MM_MX_A16F4数据流**:
108+ - `weight`数据类型必须为`torch_npu.float4_e2m1fn_x2`,使用`torch.uint8`承载(1个`torch.uint8`元素打包2个fp4数据),数据格式为$ND$,支持2维输入,转置与非转置均支持。
109+ - 非转置时`weight`的shape为$\{K, N/2\}$(沿N维打包);转置时shape为$\{K/2, N\}$(沿K维打包),stride为$[1, K/2]$;打包维度的元素个数须为偶数,且K必须满足$K \% 32 = 0$。
110+ - `weight_scale`数据类型必须为`torch.float8_e8m0fnu`,shape为$\{ceil\_div(K, 32), N\}$;`weight`转置输入时,`weight_scale`需与`weight`布局一致(同为转置视图,stride为$[1, ceil\_div(K, 32)]$)。
111+ - `k_group_size`必须等于`32`。
112+ - `x_dtype`必须为`torch.float16`或`torch.bfloat16`。
113+ - `x_scale_dtype`必须传入`None`。
114+ - 当前不支持`weight_offset`,必须传入`None`。
95 115 
96## 调用示例116## 调用示例
97 117 
@@ -121,7 +141,7 @@ torch_npu.npu_weight_quant_preprocess(weight, weight_scale, x_dtype, weight_dtyp
121 weight,141 weight,
122 weight_scale,142 weight_scale,
123 x_dtype=torch.float8_e4m3fn,143 x_dtype=torch.float8_e4m3fn,
124- weight_dtype=torch.float4_e2m1fn_x2,144+ weight_dtype=torch_npu.float4_e2m1fn_x2,
125 weight_scale_dtype=torch.float8_e8m0fnu,145 weight_scale_dtype=torch.float8_e8m0fnu,
126 weight_offset=None,146 weight_offset=None,
127 bias=None,147 bias=None,
@@ -157,7 +177,7 @@ torch_npu.npu_weight_quant_preprocess(weight, weight_scale, x_dtype, weight_dtyp
157 weight,177 weight,
158 weight_scale,178 weight_scale,
159 x_dtype=torch.float8_e4m3fn,179 x_dtype=torch.float8_e4m3fn,
160- weight_dtype=torch.float4_e2m1fn_x2,180+ weight_dtype=torch_npu.float4_e2m1fn_x2,
161 weight_scale_dtype=torch.float8_e8m0fnu,181 weight_scale_dtype=torch.float8_e8m0fnu,
162 weight_offset=None,182 weight_offset=None,
163 bias=None,183 bias=None,
@@ -215,3 +235,68 @@ torch_npu.npu_weight_quant_preprocess(weight, weight_scale, x_dtype, weight_dtyp
215 k_group_size=k_group_size235 k_group_size=k_group_size
216 )236 )
217 ```237 ```
238+ 
239+- MM_A16F4数据流场景(pergroup、`weight`非转置,输出$FRACTAL\_NZ\_C0\_8$)
240+ 
241+ ```python
242+ import torch
243+ import torch_npu
244+ 
245+ # MM_A16F4 数据流示例
246+ k, n = 256, 128
247+ k_group_size = 64
248+ g = (k + k_group_size - 1) // k_group_size
249+ 
250+ # weight: float4_e2m1,每个 uint8 打包2个fp4数据
251+ # 非转置输入:沿N维打包,logical shape {K, N},packed shape {K, N/2}
252+ weight = torch.randint(0, 255, (k, n // 2), dtype=torch.uint8).npu()
253+ 
254+ # weight_scale: float16,pergroup shape {G, N}
255+ weight_scale = torch.randn((g, n), dtype=torch.float16).npu()
256+ 
257+ out_weight, out_weight_scale, out_weight_offset, out_bias = torch_npu.npu_weight_quant_preprocess(
258+ weight,
259+ weight_scale,
260+ x_dtype=torch.float16,
261+ weight_dtype=torch_npu.float4_e2m1fn_x2,
262+ weight_scale_dtype=torch.float16,
263+ k_group_size=k_group_size
264+ )
265+ ```
266+ 
267+- MM_MX_A16F4数据流场景
268+ 
269+ ```python
270+ import torch
271+ import torch_npu
272+ 
273+ # MM_MX_A16F4 数据流示例
274+ # 开关:weight 转置状态
275+ weight_trans = False # True 表示 weight 转置输入
276+ 
277+ k, n = 256, 128
278+ g = k // 32 # k_group_size=32,scale 的 K 维组数
279+ 
280+ # weight: float4_e2m1,每个 uint8 打包2个fp4数据
281+ # weight_scale: float8_e8m0,使用 uint8 承载,shape {K/32, N}
282+ if weight_trans:
283+ # 转置输入:weight 沿K维打包,view shape {K/2, N},stride [1, K/2]
284+ # 转置输入时 weight_scale 需与 weight 布局一致(同为转置视图)
285+ weight = torch.randint(0, 255, (n, k // 2), dtype=torch.uint8).npu().transpose(0, 1)
286+ weight_scale = torch.randint(0, 255, (n, g), dtype=torch.uint8).view(
287+ torch.float8_e8m0fnu).npu().transpose(0, 1)
288+ else:
289+ # 非转置输入:weight 沿N维打包,shape {K, N/2}
290+ weight = torch.randint(0, 255, (k, n // 2), dtype=torch.uint8).npu()
291+ weight_scale = torch.randint(0, 255, (g, n), dtype=torch.uint8).view(
292+ torch.float8_e8m0fnu).npu()
293+ 
294+ out_weight, out_weight_scale, out_weight_offset, out_bias = torch_npu.npu_weight_quant_preprocess(
295+ weight,
296+ weight_scale,
297+ x_dtype=torch.float16,
298+ weight_dtype=torch_npu.float4_e2m1fn_x2,
299+ weight_scale_dtype=torch.float8_e8m0fnu,
300+ k_group_size=32
301+ )
302+ ```
@@ -51,33 +51,21 @@ at::Tensor npu_weight_quant_batchmatmul(
51 OPS_ERROR(ErrCode::PARAM));51 OPS_ERROR(ErrCode::PARAM));
52 auto x_k_dim = x.size(x_dim_num - 1);52 auto x_k_dim = x.size(x_dim_num - 1);
53 53 
54- // 计算 weight 的 K 维度54+ // 打包维还原 = 载体位宽 × 打包方向:int32/float 载体打包 8 个 4-bit,uint8 载体
55- // 对于 INT32/FLOAT 类型的 4-bit 打包(INT4_NUMS_IN_INT32 = 8)55+ // (A16W4 链路)打包 2 个;转置视图沿 K 打包(K 按 ratio 还原),否则沿最后一维 N
56- // uint8 载体 4-bit 紧凑排布(A16S4 链路):NZ_C0_16 的 view 是逻辑 [K, N],NZ_C0_8 的 view 是56+ // 打包(N 按 ratio 还原)。NZ_C0_8 的 view 本身即连续物理打包形状 [K, N/2],同规则覆盖
57- // 物理打包形状 [K, N/2];ND 的 view 是物理打包形状——非转置 [K, N/2](沿 N 打包,行主序连续,
58- // stride(-1)==1 且 stride(-2)==size(-1))、转置 [K/2, N](沿 K 打包)。K=2 时转置视图退化为
59- // [1, N] strides [1, 1],需靠 stride(-2)==size(-1) 排除误判,K/N 需按打包方向还原
60 bool is_int32_float_packed = (weight.dtype() == at::kInt || weight.dtype() == at::kFloat);57 bool is_int32_float_packed = (weight.dtype() == at::kInt || weight.dtype() == at::kFloat);
61- 
62- int64_t weight_format = at_npu::native::custom_ops::get_npu_format(weight);
63 aclDataType weight_acl_dtype =58 aclDataType weight_acl_dtype =
64 weight_dtype.has_value() ? c10_npu::GetAclDataType(weight_dtype.value()) : ACL_DT_UNDEFINED;59 weight_dtype.has_value() ? c10_npu::GetAclDataType(weight_dtype.value()) : ACL_DT_UNDEFINED;
65 bool is_4bit_acl_dtype =60 bool is_4bit_acl_dtype =
66 (weight_acl_dtype == ACL_INT4 || weight_acl_dtype == ACL_FLOAT4_E2M1 || weight_acl_dtype == ACL_FLOAT4_E1M2);61 (weight_acl_dtype == ACL_INT4 || weight_acl_dtype == ACL_FLOAT4_E2M1 || weight_acl_dtype == ACL_FLOAT4_E1M2);
67- bool is_uint8_4bit_nd = (weight.dtype() == at::kByte) && is_4bit_acl_dtype && (weight_format == ACL_FORMAT_ND);62+ int64_t weight_format = at_npu::native::custom_ops::get_npu_format(weight);
68- bool is_uint8_4bit_nz_c08 =63+ // at::kByte 即 uint8(ATen 无 kUInt8 别名,Byte 是 uint8 的唯一拼写)
69- (weight.dtype() == at::kByte) && is_4bit_acl_dtype && (weight_format == ACL_FORMAT_FRACTAL_NZ_C0_8);64+ bool is_uint8_4bit = (weight.dtype() == at::kByte) && is_4bit_acl_dtype;
X
Xxubinglin23 天前

如果有uint8的话替换kByte会更加清晰

likedislike
马琦钧
马琦钧
23 天前 评论:
70- bool uint8_pack_along_n = is_uint8_4bit_nd &&
71- (weight.stride(weight_dim_num - 1) == 1 &&
72- weight.stride(weight_dim_num - MINIMUM_SHAPE_SIZE) == weight.size(weight_dim_num - 1));
73- bool uint8_pack_along_k = is_uint8_4bit_nd && !uint8_pack_along_n;
74 65 
75- int64_t weight_k_dim = weight.size(weight_dim_num - MINIMUM_SHAPE_SIZE);66+ int64_t pack_ratio = is_int32_float_packed ? INT4_NUMS_IN_INT32 : (is_uint8_4bit ? B4_IN_UINT8 : 1);
76- if (is_int32_float_packed && trans_weight) {67+ 
77- weight_k_dim *= INT4_NUMS_IN_INT32; // 868+ int64_t weight_k_dim = weight.size(weight_dim_num - MINIMUM_SHAPE_SIZE) * (trans_weight ? pack_ratio : 1);
78- } else if (uint8_pack_along_k) {
79- weight_k_dim *= B4_IN_UINT8; // 沿 K 打包,K 按每字节 2 个 4-bit 还原
80- }
81 69 
82 TORCH_CHECK(70 TORCH_CHECK(
83 x_k_dim == weight_k_dim,71 x_k_dim == weight_k_dim,
@@ -92,13 +80,7 @@ at::Tensor npu_weight_quant_batchmatmul(
92 output_size.resize(out_dim_num, 1);80 output_size.resize(out_dim_num, 1);
93 output_size[out_dim_num - MINIMUM_SHAPE_SIZE] = x.size(x_dim_num - MINIMUM_SHAPE_SIZE);81 output_size[out_dim_num - MINIMUM_SHAPE_SIZE] = x.size(x_dim_num - MINIMUM_SHAPE_SIZE);
94 auto weight_size_base = weight.size(weight_dim_num - MINIMUM_SHAPE_SIZE + 1);82 auto weight_size_base = weight.size(weight_dim_num - MINIMUM_SHAPE_SIZE + 1);
95- if (is_int32_float_packed && !trans_weight) {83+ output_size[out_dim_num - MINIMUM_SHAPE_SIZE + 1] = weight_size_base * (trans_weight ? 1 : pack_ratio);
96- output_size[out_dim_num - MINIMUM_SHAPE_SIZE + 1] = weight_size_base * INT4_NUMS_IN_INT32;
97- } else if (uint8_pack_along_n || is_uint8_4bit_nz_c08) {
98- output_size[out_dim_num - MINIMUM_SHAPE_SIZE + 1] = weight_size_base * B4_IN_UINT8;
99- } else {
100- output_size[out_dim_num - MINIMUM_SHAPE_SIZE + 1] = weight_size_base;
101- }
102 if (x_dim_num == weight_dim_num) {84 if (x_dim_num == weight_dim_num) {
103 for (auto i = 0; i < out_dim_num - MINIMUM_SHAPE_SIZE; i++) {85 for (auto i = 0; i < out_dim_num - MINIMUM_SHAPE_SIZE; i++) {
104 TORCH_CHECK(x.size(i) == weight.size(i), "batch of x is diff from batch of weight", OPS_ERROR(ErrCode::PARAM));86 TORCH_CHECK(x.size(i) == weight.size(i), "batch of x is diff from batch of weight", OPS_ERROR(ErrCode::PARAM));
@@ -77,7 +77,7 @@ static bool is_transpose_certain_two_dims(const at::Tensor& tensor, int64_t firs
77 return tensor.stride(first_dim + 1) == tensor.stride(first_dim) * tensor.size(first_dim);77 return tensor.stride(first_dim + 1) == tensor.stride(first_dim) * tensor.size(first_dim);
78}78}
79 79 
80-// A16W4(INT4 / FP4 E2M1,非 MX)torch 侧统一用 uint8 载体物理打包(每字节 2 个 4-bit 元素),80+// A16W4(INT4 / FP4 E2M1,含 MX)torch 侧统一用 uint8 载体物理打包(每字节 2 个 4-bit 元素),
81// 与 aclnnWeightQuantBatchMatmulV2 的 uint8 packed 输入约定一致;其他载体直接拒绝81// 与 aclnnWeightQuantBatchMatmulV2 的 uint8 packed 输入约定一致;其他载体直接拒绝
82static void check_a16w4_uint8_carrier(const QuantContext& ctx) {82static void check_a16w4_uint8_carrier(const QuantContext& ctx) {
83 TORCH_CHECK(83 TORCH_CHECK(
@@ -200,7 +200,7 @@ bool judge_mm_a16s4_per_channel(QuantContext& ctx) {
200 if (x_dtype_match && weight_acl_dtype == ACL_INT4 && ctx.weight.dim() == DIMS_2) {200 if (x_dtype_match && weight_acl_dtype == ACL_INT4 && ctx.weight.dim() == DIMS_2) {
201 int64_t scale_dim = ctx.weight_scale.dim();201 int64_t scale_dim = ctx.weight_scale.dim();
202 bool is_per_channel = (scale_dim == DIMS_1) || (scale_dim == DIMS_2 && ctx.weight_scale.size(0) == 1);202 bool is_per_channel = (scale_dim == DIMS_1) || (scale_dim == DIMS_2 && ctx.weight_scale.size(0) == 1);
203- // 转置状态内部分流:转置 → ND 直拷,非转置 → NZ 转换(prepare_out_weight_a16s4)203+ // 转置状态内部分流:转置 → ND 直拷,非转置 → NZ 转换(prepare_out_weight_a16w4)
204 if (is_per_channel && ctx.weight_scale.numel() > 1) {204 if (is_per_channel && ctx.weight_scale.numel() > 1) {
205 check_a16w4_uint8_carrier(ctx);205 check_a16w4_uint8_carrier(ctx);
206 ctx.is_weight_trans = is_transpose_certain_two_dims(ctx.weight, 0);206 ctx.is_weight_trans = is_transpose_certain_two_dims(ctx.weight, 0);
@@ -218,7 +218,7 @@ bool judge_mm_a16s4_per_group(QuantContext& ctx) {
218 218 
219 bool x_dtype_match = (x_acl_dtype == ACL_FLOAT16 || x_acl_dtype == ACL_BF16);219 bool x_dtype_match = (x_acl_dtype == ACL_FLOAT16 || x_acl_dtype == ACL_BF16);
220 220 
221- // A16S4 per-group:转置状态内部分流,转置 → ND 直拷,非转置 → NZ 转换(prepare_out_weight_a16s4)221+ // A16S4 per-group:转置状态内部分流,转置 → ND 直拷,非转置 → NZ 转换(prepare_out_weight_a16w4)
222 if (x_dtype_match && weight_acl_dtype == ACL_INT4 && x_scale_acl_dtype == ACL_DT_UNDEFINED &&222 if (x_dtype_match && weight_acl_dtype == ACL_INT4 && x_scale_acl_dtype == ACL_DT_UNDEFINED &&
223 ctx.weight.dim() == DIMS_2 && ctx.weight_scale.dim() == DIMS_2 && ctx.weight_scale.size(0) > 1) {223 ctx.weight.dim() == DIMS_2 && ctx.weight_scale.dim() == DIMS_2 && ctx.weight_scale.size(0) > 1) {
224 check_a16w4_uint8_carrier(ctx);224 check_a16w4_uint8_carrier(ctx);
@@ -239,6 +239,7 @@ bool judge_mm_a16f4_nz_pergroup(QuantContext& ctx) {
239 bool scale_dtype_match = (weight_scale_acl_dtype == ACL_FLOAT16 || weight_scale_acl_dtype == ACL_BF16);239 bool scale_dtype_match = (weight_scale_acl_dtype == ACL_FLOAT16 || weight_scale_acl_dtype == ACL_BF16);
240 240 
241 // A16F4 per-group NZ:FP4 weight + per-group scale [G, N](G > 1),仅支持非转置 weight241 // A16F4 per-group NZ:FP4 weight + per-group scale [G, N](G > 1),仅支持非转置 weight
242+ // (pergroup 转置无下游 wqbmmv2 支持,整体不支持;MX 转置 ND 直拷见 judge_mm_a16f4_mx)
242 if (x_dtype_match && scale_dtype_match && weight_acl_dtype == ACL_FLOAT4_E2M1 &&243 if (x_dtype_match && scale_dtype_match && weight_acl_dtype == ACL_FLOAT4_E2M1 &&
243 x_scale_acl_dtype == ACL_DT_UNDEFINED && ctx.weight.dim() == DIMS_2 && ctx.weight_scale.dim() == DIMS_2 &&244 x_scale_acl_dtype == ACL_DT_UNDEFINED && ctx.weight.dim() == DIMS_2 && ctx.weight_scale.dim() == DIMS_2 &&
244 ctx.weight_scale.size(0) > 1 && !is_transpose_certain_two_dims(ctx.weight, 0)) {245 ctx.weight_scale.size(0) > 1 && !is_transpose_certain_two_dims(ctx.weight, 0)) {
@@ -258,44 +259,25 @@ bool judge_mm_a16f4_mx(QuantContext& ctx) {
258 259 
259 bool x_dtype_match = (x_acl_dtype == ACL_FLOAT16 || x_acl_dtype == ACL_BF16);260 bool x_dtype_match = (x_acl_dtype == ACL_FLOAT16 || x_acl_dtype == ACL_BF16);
260 261 
261- // A16 MXFP4:FP4 weight + MX scale(E8M0,2D [K/32, N] 连续),仅支持非转置 weight262+ // A16 MXFP4:FP4 weight + MX scale(E8M0,2D [K/32, N],转置时为 strides [1, K/32] 视图);
263+ // 转置状态内部分流:转置 → ND 直拷(wqbmmv2 MX kernel 支持 ND 转置),非转置 → NZ 转换
262 if (x_dtype_match && weight_acl_dtype == ACL_FLOAT4_E2M1 && weight_scale_acl_dtype == ACL_FLOAT8_E8M0 &&264 if (x_dtype_match && weight_acl_dtype == ACL_FLOAT4_E2M1 && weight_scale_acl_dtype == ACL_FLOAT8_E8M0 &&
263- x_scale_acl_dtype == ACL_DT_UNDEFINED && ctx.weight.dim() == DIMS_2 && ctx.weight_scale.dim() == DIMS_2 &&265+ x_scale_acl_dtype == ACL_DT_UNDEFINED && ctx.weight.dim() == DIMS_2 && ctx.weight_scale.dim() == DIMS_2) {
264- !is_transpose_certain_two_dims(ctx.weight, 0)) {
265 check_a16w4_uint8_carrier(ctx);266 check_a16w4_uint8_carrier(ctx);
266- ctx.is_weight_trans = false;267+ ctx.is_weight_trans = is_transpose_certain_two_dims(ctx.weight, 0);
267 return true;268 return true;
268 }269 }
269 return false;270 return false;
270}271}
271 272 
272-// 校验镜像 strides 的寻址包络不超出按 numel 最小分配的 buffer;273+// 直拷输出别名:out 直接复用输入 tensor(共享 storage 与 view),aclnn 直拷路径同址时无拷贝
273-// 连续/确定转置视图必过,padding/重叠/expand 视图直接拒绝而不是写出界274+static void prepare_out_weight_scale_direct(QuantContext& ctx) {
274-static void check_strides_envelope(c10::IntArrayRef sizes, c10::IntArrayRef strides, int64_t numel) {275+ ctx.out_weight_scale = ctx.weight_scale;
275- int64_t max_offset = 0;
276- for (size_t i = 0; i < sizes.size(); ++i) {
277- TORCH_CHECK(strides[i] > 0, "expanded/overlapped view is not supported here", OPS_ERROR(ErrCode::PARAM));
278- max_offset += (sizes[i] - 1) * strides[i];
279- }
280- TORCH_CHECK(max_offset < numel, "mirrored strides exceed output allocation: max offset ", max_offset,
281- " vs numel ", numel, OPS_ERROR(ErrCode::PARAM));
282-}
283- 
284-static void prepare_out_weight_scale(QuantContext& ctx) {
285- auto scale_view_shape = op_infer::array_to_small_vector(ctx.weight_scale.sizes());
286- ctx.out_weight_scale = npu_preparation::apply_tensor_without_format(scale_view_shape, ctx.weight_scale.options());
287- check_strides_envelope(ctx.weight_scale.sizes(), ctx.weight_scale.strides(), ctx.out_weight_scale.numel());
288- // 保持与输入 scale 相同的 strides(ND per-group 转置场景需配转置 scale),连续输入时为空操作
289- ctx.out_weight_scale =
290- ctx.out_weight_scale.as_strided_(scale_view_shape, op_infer::array_to_small_vector(ctx.weight_scale.strides()));
291}276}
292 277 
293static void prepare_out_weight_nd(QuantContext& ctx) {278static void prepare_out_weight_nd(QuantContext& ctx) {
294- auto weight_view_shape = op_infer::array_to_small_vector(ctx.weight.sizes());279+ // ND 直拷为物理透传,out_weight 直接别名 weight(共享 storage 与 view,转置场景尤为关键)
295- ctx.out_weight = npu_preparation::apply_tensor_without_format(weight_view_shape, ctx.weight.options());280+ ctx.out_weight = ctx.weight;
296- check_strides_envelope(ctx.weight.sizes(), ctx.weight.strides(), ctx.out_weight.numel());
297- // ND 直拷是物理透传,out_weight 必须与 weight 保持相同的 sizes/strides(转置场景尤为关键)
298- ctx.out_weight = ctx.out_weight.as_strided_(weight_view_shape, op_infer::array_to_small_vector(ctx.weight.strides()));
299}281}
300 282 
301template <bool IsGmm, int64_t NzC0, aclFormat OutWeightFormat>283template <bool IsGmm, int64_t NzC0, aclFormat OutWeightFormat>
@@ -377,8 +359,8 @@ static void prepare_out_weight_nz_a16w4(QuantContext& ctx) {
377 ctx.out_weight, ctx.out_weight.sizes(), storage_shape, ctx.out_weight.strides(), ACL_FORMAT_FRACTAL_NZ_C0_8);359 ctx.out_weight, ctx.out_weight.sizes(), storage_shape, ctx.out_weight.strides(), ACL_FORMAT_FRACTAL_NZ_C0_8);
378}360}
379 361 
380-// A16S4 内部分流:转置 weight → ND 直拷(物理透传,与转置状态无关);非转置 → NZ 转换362+// A16W4(INT4/FP4)转置分流:转置 weight → ND 直拷(物理透传,与转置布局无关);非转置 → NZ 转换
381-static void prepare_out_weight_a16s4(QuantContext& ctx) {363+static void prepare_out_weight_a16w4(QuantContext& ctx) {
382 if (ctx.is_weight_trans) {364 if (ctx.is_weight_trans) {
383 prepare_out_weight_nd(ctx);365 prepare_out_weight_nd(ctx);
384 } else {366 } else {
@@ -418,24 +400,17 @@ static void prepare_out_weight_scale_mx(QuantContext& ctx) {
418 }400 }
419}401}
420 402 
421-static void prepare_out_weight_offset(QuantContext& ctx) {403+static void prepare_out_weight_offset_direct(QuantContext& ctx) {
422 if (ctx.weight_offset.has_value() && ctx.weight_offset.value().defined()) {404 if (ctx.weight_offset.has_value() && ctx.weight_offset.value().defined()) {
423- auto offset_sizes = op_infer::array_to_small_vector(ctx.weight_offset.value().sizes());405+ // offset 为直拷,out 直接别名输入(共享 storage 与 view,转置场景 strides 一并继承)
424- ctx.out_weight_offset =406+ ctx.out_weight_offset = ctx.weight_offset.value();
425- npu_preparation::apply_tensor_without_format(offset_sizes, ctx.weight_offset.value().options());
426- check_strides_envelope(
427- ctx.weight_offset.value().sizes(), ctx.weight_offset.value().strides(), ctx.out_weight_offset.numel());
428- // 与 prepare_out_weight_scale 同理:保持与输入 offset 相同的 strides,
429- // 转置场景下游 wqbmmv2 要求 antiquantOffset 与 weight 连续/转置状态一致,连续输入时为空操作
430- ctx.out_weight_offset = ctx.out_weight_offset.as_strided_(
431- offset_sizes, op_infer::array_to_small_vector(ctx.weight_offset.value().strides()));
432 }407 }
433}408}
434 409 
435-static void prepare_out_bias(QuantContext& ctx) {410+static void prepare_out_bias_direct(QuantContext& ctx) {
436 if (ctx.bias.has_value() && ctx.bias.value().defined()) {411 if (ctx.bias.has_value() && ctx.bias.value().defined()) {
437- auto bias_sizes = op_infer::array_to_small_vector(ctx.bias.value().sizes());412+ // bias 为直拷,out 直接别名输入
438- ctx.out_bias = npu_preparation::apply_tensor_without_format(bias_sizes, ctx.bias.value().options());413+ ctx.out_bias = ctx.bias.value();
439 }414 }
440}415}
441 416 
@@ -451,39 +426,40 @@ static const std::unordered_map<c10_npu::SocVersion, std::vector<DataFlowConfig>
451 {// FP4 逻辑 C0 为 32,torch 层用 1 个 int8/uint8 元素打包 2 个 FP4,因此物理存储的 C0 维长度为 16426 {// FP4 逻辑 C0 为 32,torch 层用 1 个 int8/uint8 元素打包 2 个 FP4,因此物理存储的 C0 维长度为 16
452 prepare_out_weight_nz<false, NZ_C0_16, ACL_FORMAT_FRACTAL_NZ_C0_16>,427 prepare_out_weight_nz<false, NZ_C0_16, ACL_FORMAT_FRACTAL_NZ_C0_16>,
453 prepare_out_weight_scale_mx<false>,428 prepare_out_weight_scale_mx<false>,
454- prepare_out_weight_offset,429+ prepare_out_weight_offset_direct,
455- prepare_out_bias}},430+ prepare_out_bias_direct}},
456 {judge_gmm_mx_a8w4,431 {judge_gmm_mx_a8w4,
457 {// FP4 逻辑 C0 为 32,torch 层用 1 个 int8/uint8 元素打包 2 个 FP4,因此物理存储的 C0 维长度为 16432 {// FP4 逻辑 C0 为 32,torch 层用 1 个 int8/uint8 元素打包 2 个 FP4,因此物理存储的 C0 维长度为 16
458 prepare_out_weight_nz<true, NZ_C0_16, ACL_FORMAT_FRACTAL_NZ_C0_16>,433 prepare_out_weight_nz<true, NZ_C0_16, ACL_FORMAT_FRACTAL_NZ_C0_16>,
459 prepare_out_weight_scale_mx<true>,434 prepare_out_weight_scale_mx<true>,
460- prepare_out_weight_offset,435+ prepare_out_weight_offset_direct,
461- prepare_out_bias}},436+ prepare_out_bias_direct}},
462 {judge_mm_a16s4_per_tensor,437 {judge_mm_a16s4_per_tensor,
463 {prepare_out_weight_nd,438 {prepare_out_weight_nd,
464- prepare_out_weight_scale,439+ prepare_out_weight_scale_direct,
465- prepare_out_weight_offset,440+ prepare_out_weight_offset_direct,
466- prepare_out_bias}},441+ prepare_out_bias_direct}},
467 {judge_mm_a16s4_per_channel,442 {judge_mm_a16s4_per_channel,
468- {prepare_out_weight_a16s4,443+ {prepare_out_weight_a16w4,
469- prepare_out_weight_scale,444+ prepare_out_weight_scale_direct,
470- prepare_out_weight_offset,445+ prepare_out_weight_offset_direct,
471- prepare_out_bias}},446+ prepare_out_bias_direct}},
472 {judge_mm_a16s4_per_group,447 {judge_mm_a16s4_per_group,
473- {prepare_out_weight_a16s4,448+ {prepare_out_weight_a16w4,
474- prepare_out_weight_scale,449+ prepare_out_weight_scale_direct,
475- prepare_out_weight_offset,450+ prepare_out_weight_offset_direct,
476- prepare_out_bias}},451+ prepare_out_bias_direct}},
477 {judge_mm_a16f4_nz_pergroup,452 {judge_mm_a16f4_nz_pergroup,
478 {prepare_out_weight_nz_a16w4,453 {prepare_out_weight_nz_a16w4,
479- prepare_out_weight_scale,454+ prepare_out_weight_scale_direct,
480- prepare_out_weight_offset,455+ prepare_out_weight_offset_direct,
481- prepare_out_bias}},456+ prepare_out_bias_direct}},
482 {judge_mm_a16f4_mx,457 {judge_mm_a16f4_mx,
483- {prepare_out_weight_nz_a16w4,458+ {// 转置 → ND 直拷,非转置 → NZ 转换(prepare_out_weight_a16w4 内部分流)
484- prepare_out_weight_scale,459+ prepare_out_weight_a16w4,
485- prepare_out_weight_offset,460+ prepare_out_weight_scale_direct,
486- prepare_out_bias}}}},461+ prepare_out_weight_offset_direct,
462+ prepare_out_bias_direct}}}},
487};463};
488 464 
489} // namespace465} // namespace
@@ -34,6 +34,36 @@ class TestNpuWeightQuantPreprocess(TestCase):
34 if torch_npu._C._npu_getOption("ALLOW_INTERNAL_FORMAT") == b"enable":34 if torch_npu._C._npu_getOption("ALLOW_INTERNAL_FORMAT") == b"enable":
35 self.assertEqual(torch_npu.get_npu_format(out_weight), FRACTAL_NZ_C0_16)35 self.assertEqual(torch_npu.get_npu_format(out_weight), FRACTAL_NZ_C0_16)
36 36 
37+ @SupportedDevices(['Ascend950'])
38+ def test_npu_weight_quant_preprocess_a16mxf4_trans_nd(self):
39+ # A16MXF4 转置:weight 物理 {N, K/2} 沿 K 打包、视图 {K/2, N} strides [1, K/2];
40+ # scale E8M0 转置视图 {G, N} strides [1, G](G=K/32)。转置走 ND 直拷物理透传,
41+ # 输出与输入共享内容,format 保持 ND
42+ k, n = 256, 128
43+ g = k // 32
44+ weight = torch.zeros((n, k // 2), dtype=torch.uint8).npu().transpose(0, 1)
45+ weight_scale = torch.zeros((n, g), dtype=torch.uint8).view(
46+ torch.float8_e8m0fnu).npu().transpose(0, 1)
47+ 
48+ out_weight, out_weight_scale, _, _ = torch_npu.npu_weight_quant_preprocess(
49+ weight,
50+ weight_scale,
51+ x_dtype=torch.float16,
52+ weight_dtype=torch_npu.float4_e2m1fn_x2,
53+ weight_scale_dtype=torch_npu.float8_e8m0fnu,
54+ k_group_size=32)
55+ 
56+ self.assertEqual(out_weight.shape, weight.shape)
57+ self.assertEqual(out_weight.stride(), weight.stride())
58+ self.assertEqual(out_weight.dtype, weight.dtype)
59+ self.assertEqual(out_weight_scale.shape, (g, n))
60+ self.assertEqual(out_weight_scale.stride(), weight_scale.stride())
61+ self.assertEqual(out_weight.cpu(), weight.cpu())
62+ # e8m0 转置视图直接 .cpu() 会走 AICPU Transpose(不支持该 dtype),按 uint8 回读比对
63+ self.assertEqual(out_weight_scale.view(torch.uint8).cpu(), weight_scale.view(torch.uint8).cpu())
64+ # ND 直拷透传:out_weight format 与输入一致(ND;返回值类型随 ALLOW_INTERNAL_FORMAT 开关变化)
65+ self.assertEqual(torch_npu.get_npu_format(out_weight), torch_npu.get_npu_format(weight))
66+ 
37 @SupportedDevices(['Ascend950'])67 @SupportedDevices(['Ascend950'])
38 def test_npu_weight_quant_preprocess_a16w4_int8_carrier_rejected(self):68 def test_npu_weight_quant_preprocess_a16w4_int8_carrier_rejected(self):
39 # A16W4 的 4-bit weight 只接受 uint8 载体(每字节 2 个 4-bit),int8 载体应被拒绝69 # A16W4 的 4-bit weight 只接受 uint8 载体(每字节 2 个 4-bit),int8 载体应被拒绝