torch_npu.npu_quantize
产品支持情况
| 产品 | 是否支持 |
|---|---|
| Atlas A3 训练系列产品/Atlas A3 推理系列产品 | √ |
| Atlas A2 训练系列产品/Atlas A2 推理系列产品 | √ |
| Atlas 推理系列产品 | √ |
功能说明
- API功能:对输入的张量进行量化处理。
- 计算公式:
-
若
div_mode为True:result=(input/scales)+zero_pointsresult=(input/scales)+zero\_points
-
若
div_mode为False:result=(input∗scales)+zero_pointsresult=(input*scales)+zero\_points
-
函数原型
torch_npu.npu_quantize(input, scales, zero_points, dtype, axis=1, div_mode=True) -> Tensor
参数说明
-
input (
Tensor):必选参数,需要进行量化的源数据张量,对应公式中的input。数据格式支持NDND,支持空Tensor,支持非连续的Tensor。div_mode为False且dtype为quint4x2时,最后一维需要能被8整除。- Atlas 推理系列产品:数据类型支持
float、float16。 - Atlas A2 训练系列产品/Atlas A2 推理系列产品、Atlas A3 训练系列产品/Atlas A3 推理系列产品:数据类型支持
float、float16、bfloat16。div_mode为False时,且当数据类型为float时,数据格式支持NZNZ。
- Atlas 推理系列产品:数据类型支持
-
scales (
Tensor):必选参数,对input进行缩放的张量,对应公式中的scales。支持空Tensor,支持非连续的Tensor。-
div_mode为True时:- Atlas 推理系列产品:数据类型支持
float。 - Atlas A2 训练系列产品/Atlas A2 推理系列产品、Atlas A3 训练系列产品/Atlas A3 推理系列产品:数据类型支持
float、bfloat16。
- Atlas 推理系列产品:数据类型支持
-
div_mode为False时,数据格式支持NDND。支持1维或多维(1维时,对应轴的大小需要与input中第axis维相等或等于1;多维时,scales的shape需要与input的shape维度相等,除axis指定的维度,其他维度为1,axis指定的维度必须和input对应的维度相等或等于1)。数据类型、数据格式需要和input的数据类型和数据格式一致。- Atlas 推理系列产品:数据类型支持
float、float16。 - Atlas A2 训练系列产品/Atlas A2 推理系列产品、Atlas A3 训练系列产品/Atlas A3 推理系列产品:数据类型支持
float、float16、bfloat16。当数据格式为NZNZ时,scales的所有元素值为1。
- Atlas 推理系列产品:数据类型支持
-
-
zero_points (
Tensor):必选参数,允许为None,对input进行偏移的张量,对应公式中的zero_points。支持空Tensor,支持非连续的Tensor。-
div_mode为True时- Atlas 推理系列产品:数据类型支持
int8、uint8、int32。 - Atlas A2 训练系列产品/Atlas A2 推理系列产品、Atlas A3 训练系列产品/Atlas A3 推理系列产品:数据类型支持
int8、uint8、int32、bfloat16。
- Atlas 推理系列产品:数据类型支持
-
div_mode为False时,数据格式支持NDND。支持1维或多维(1维时,对应轴的大小需要与input中第axis维相等或等于1;多维时,scales的shape需要与input维度相等,除axis指定的维度,其他维度为1,axis指定的维度必须和input对应的维度相等)。zero_points的shape和dtype需要和scales一致。- Atlas 推理系列产品:数据类型支持
float、float16。 - Atlas A2 训练系列产品/Atlas A2 推理系列产品、Atlas A3 训练系列产品/Atlas A3 推理系列产品:数据类型支持
float、float16、bfloat16。当input的数据格式为NZNZ时,值为空。
- Atlas 推理系列产品:数据类型支持
-
-
dtype (
int):必选参数,指定输出参数的类型。-
div_mode为True时,- Atlas 推理系列产品:类型支持
qint8、quint8、int32。 - Atlas A2 训练系列产品/Atlas A2 推理系列产品、Atlas A3 训练系列产品/Atlas A3 推理系列产品:类型支持
qint8、quint8、int32。
- Atlas 推理系列产品:类型支持
-
div_mode为False时,类型支持qint8、quint4x2。如果dtype为quint4x2时,输出tensor类型为int32,由8个int4拼接。
-
-
axis (
int):可选参数,量化的element-wise轴,其他的轴做broadcast,默认值为1。div_mode为False时,axis取值范围是[-2, +∞)且指定的轴不能超过输入input的维度数。如果axis为-2,代表量化的element-wise轴是输入input的倒数第二根轴;如果axis大于-2,量化的element-wise轴是输入的最后一根轴。 -
div_mode (
bool):可选参数,表示计算scales模式,对应公式中的div_mode。当div_mode为True时,表示用除法计算scales;div_mode为False时,表示用乘法计算scales,默认值为True。
返回值说明
Tensor
对应公式中的result。数据类型由参数dtype指定,如果参数dtype为quint4x2,输出的dtype是int32,shape的最后一维是input的shape最后一维的1/8,shape其他维度和input的shape其他维度保持一致;如果参数dtype不为quint4x2时,shape与输入input的shape保持一致。输出的数据格式与输入input的数据格式保持一致,且当数据格式为NZNZ时,数据类型仅支持INT32。支持空Tensor,支持非连续的Tensor。
约束说明
- 该接口支持推理场景下使用。
- 该接口支持图模式。
input数据格式为NZNZ时,input输入shape支持3维,形如(e, k, n), k为256倍数,n为8的倍数,scales输入shape支持1维或3维,zero_points输入为None,dtype为quint4x2。div_mode为False时:- 支持Atlas A2 训练系列产品/Atlas A2 推理系列产品。
- 当
dtype为quint4x2或者axis为-2时,不支持Atlas 推理系列产品。
调用示例
-
单算子模式调用
-
Atlas A2 训练系列产品/Atlas A2 推理系列产品
>>> import torch >>> import torch_npu >>> >>> x = torch.randn(1, 1, 12).bfloat16().npu() >>> scale = torch.tensor([0.1] * 12).bfloat16().npu() >>> out = torch_npu.npu_quantize(x, scale, None, torch.qint8, -1, False) >>> x tensor([[[ 0.9609, 1.3281, -0.6172, 0.5469, -1.1797, -1.1719, -0.7422, 0.9727, -0.9062, -0.0815, -0.8047, 1.0703]]], device='npu:0', dtype=torch.bfloat16) >>> out tensor([[[0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0]]], device='npu:0', dtype=torch.int8) -
Atlas 推理系列产品
>>> import torch >>> import torch_npu >>> >>> x = torch.randn((2, 3, 12), dtype=torch.float).npu() >>> scale = torch.tensor(([3] * 12),dtype=torch.float).npu() >>> out = torch_npu.npu_quantize(x, scale, None, torch.qint8, -1, False) >>> x tensor([[[-7.7834e-01, -1.0473e+00, -1.1155e+00, 1.2233e+00, -1.2271e+00, -2.5612e+00, -1.8274e-01, 2.8293e+00, 1.9029e-01, -1.9333e+00, -4.9270e-01, -1.0650e+00], [-8.9416e-01, 3.1869e-02, -5.8144e-01, -4.9477e-01, 9.7323e-02, -3.8681e-01, 2.1969e-03, -6.3244e-01, 7.1591e-01, -1.8587e-01, -1.3381e+00, -2.6253e-01], [ 1.8462e-02, 1.2397e-01, -9.0656e-01, -9.9280e-01, -4.4235e-02, 1.0623e+00, -9.8437e-02, 1.2941e+00, 1.0805e+00, -1.7269e-01, -9.9205e-02, -6.1429e-01]], [[ 1.3678e+00, -2.7348e-01, -4.1354e-01, -9.4638e-01, 4.2792e-01, 8.0462e-01, 9.3584e-01, 6.3704e-01, 1.1269e+00, -1.5329e+00, 5.8572e-01, -1.3966e+00], [ 3.5882e-01, 8.7029e-01, -1.3176e+00, 1.1601e+00, -3.6984e-01, 7.3642e-01, -1.0755e+00, 6.6557e-01, 3.1149e+00, -6.8776e-01, -1.0913e+00, 4.4962e-01], [-1.2505e+00, 1.5474e+00, -7.4332e-02, -1.6657e+00, 1.3275e+00, 5.8914e-02, 8.4287e-01, -1.7109e+00, 1.8256e-01, 3.2937e-01, 2.4875e+00, 1.3921e+00]]], device='npu:0') >>> out tensor([[[-2, -3, -3, 4, -4, -8, -1, 8, 1, -6, -1, -3], [-3, 0, -2, -1, 0, -1, 0, -2, 2, -1, -4, -1], [ 0, 0, -3, -3, 0, 3, 0, 4, 3, -1, 0, -2]], [[ 4, -1, -1, -3, 1, 2, 3, 2, 3, -5, 2, -4], [ 1, 3, -4, 3, -1, 2, -3, 2, 9, -2, -3, 1], [-4, 5, 0, -5, 4, 0, 3, -5, 1, 1, 7, 4]]], device='npu:0', dtype=torch.int8)
-
-
图模式调用
import torch import torch_npu import torchair as tng from torchair.ge_concrete_graph import ge_apis as ge from torchair.configs.compiler_config import CompilerConfig x = torch.randn((2, 3, 12), dtype=torch.float16).npu() scale = torch.tensor(([3] * 12), dtype=torch.float16).npu() axis = 1 div_mode = False class Network(torch.nn.Module): def __init__(self): super(Network, self).__init__() def forward(self, x, scale, zero_points, dst_type, div_mode): return torch_npu.npu_quantize(x, scale, zero_points=zero_points, dtype=dst_type, div_mode=div_mode) model = Network() config = CompilerConfig() npu_backend = tng.get_npu_backend(compiler_config=config) config.debug.graph_dump.type = 'pbtxt' model = torch.compile(model, fullgraph=True, backend=npu_backend, dynamic=True) output_data = model(x, scale, None, dst_type=torch.qint8, div_mode=div_mode) print("shape of x:", x.shape) print("shape of output_data:", output_data.shape) print("x:", x) print("output_data:", output_data) # 执行上述代码的输出类似如下 shape of x: torch.Size([2, 3, 12]) shape of output_data: torch.Size([2, 3, 12]) x: tensor([[[ 0.2600, 0.6782, 0.5024, 0.9492, 0.6089, 0.7461, 1.5332, -0.2123, 0.6558, -0.8354, -0.5366, -0.6821], [-0.2522, 0.2415, -0.0269, -0.1497, 0.2256, -0.5239, 0.7363, -0.2468, 1.6064, 1.4170, -0.2213, 1.5947], [-0.6328, 0.8105, 0.2532, 1.0684, -1.2119, -0.6865, 0.7451, -0.8120, 0.6401, -2.1270, -0.9482, -1.1973]], [[-1.7461, -1.1758, -0.5352, 1.5938, 1.8945, -2.2500, -0.5073, -0.8164, 0.8267, -0.4377, 1.2490, 0.2415], [ 0.8062, -1.0498, -0.8345, 1.1465, -0.7349, 0.1317, 0.2280, -0.8145, 0.2673, 1.4756, -1.6768, 1.1572], [-0.3147, -0.4446, -1.0508, 0.8325, 1.4590, 0.2096, -0.9961, 0.6089, -0.2460, 1.1543, 0.9277, 0.1079]]], device='npu:0', dtype=torch.float16) output_data: tensor([[[ 1, 2, 2, 3, 2, 2, 5, -1, 2, -3, -2, -2], [-1, 1, 0, 0, 1, -2, 2, -1, 5, 4, -1, 5], [-2, 2, 1, 3, -4, -2, 2, -2, 2, -6, -3, -4]], [[-5, -4, -2, 5, 6, -7, -2, -2, 2, -1, 4, 1], [ 2, -3, -3, 3, -2, 0, 1, -2, 1, 4, -5, 3], [-1, -1, -3, 2, 4, 1, -3, 2, -1, 3, 3, 0]]], device='npu:0', dtype=torch.int8)