已合并
fix: WeightQuantPreprocess A16MXF4 支持转置 weight ND 直拷 #5803
马琦钧创建于 8月31日
fix: WeightQuantPreprocess A16MXF4 支持转置 weight ND 直拷 #5803
已合并
共 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_size | 235 | 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 | |||
| 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; // 8 | 68 | + 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 输入约定一致;其他载体直接拒绝 |
| 82 | static void check_a16w4_uint8_carrier(const QuantContext& ctx) { | 82 | static 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),仅支持非转置 weight | 241 | // 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] 连续),仅支持非转置 weight | 262 | + // 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 | ||
| 293 | static void prepare_out_weight_nd(QuantContext& ctx) { | 278 | static 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 | ||
| 301 | template <bool IsGmm, int64_t NzC0, aclFormat OutWeightFormat> | 283 | template <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 维长度为 16 | 426 | {// 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 维长度为 16 | 432 | {// 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 | } // namespace | 465 | } // 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 | + | ||
| 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 | 67 | ||
| 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 载体应被拒绝 |
如果有uint8的话替换kByte会更加清晰