torch_npu.npu_quantize

产品支持情况

产品 是否支持
Atlas A3 训练系列产品/Atlas A3 推理系列产品
Atlas A2 训练系列产品/Atlas A2 推理系列产品
Atlas 推理系列产品

功能说明

  • API功能:对输入的张量进行量化处理。
  • 计算公式:
    • div_modeTrue

      result=(input/scales)+zero_pointsresult=(input/scales)+zero\_points

    • div_modeFalse

      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_modeFalsedtypequint4x2时,最后一维需要能被8整除。

    • Atlas 推理系列产品:数据类型支持floatfloat16
    • Atlas A2 训练系列产品/Atlas A2 推理系列产品、Atlas A3 训练系列产品/Atlas A3 推理系列产品:数据类型支持floatfloat16bfloat16div_modeFalse时,且当数据类型为float时,数据格式支持NZNZ
  • scales (Tensor):必选参数,对input进行缩放的张量,对应公式中的scales。支持空Tensor,支持非连续的Tensor。

    • div_modeTrue时:

      • Atlas 推理系列产品:数据类型支持float
      • Atlas A2 训练系列产品/Atlas A2 推理系列产品、Atlas A3 训练系列产品/Atlas A3 推理系列产品:数据类型支持floatbfloat16
    • div_modeFalse时,数据格式支持NDND。支持1维或多维(1维时,对应轴的大小需要与input中第axis维相等或等于1;多维时,scales的shape需要与input的shape维度相等,除axis指定的维度,其他维度为1,axis指定的维度必须和input对应的维度相等或等于1)。数据类型、数据格式需要和input的数据类型和数据格式一致。

      • Atlas 推理系列产品:数据类型支持floatfloat16
      • Atlas A2 训练系列产品/Atlas A2 推理系列产品、Atlas A3 训练系列产品/Atlas A3 推理系列产品:数据类型支持floatfloat16bfloat16。当数据格式为NZNZ时,scales的所有元素值为1。
  • zero_points (Tensor):必选参数,允许为None,对input进行偏移的张量,对应公式中的zero_points。支持空Tensor,支持非连续的Tensor。

    • div_modeTrue

      • Atlas 推理系列产品:数据类型支持int8uint8int32
      • Atlas A2 训练系列产品/Atlas A2 推理系列产品、Atlas A3 训练系列产品/Atlas A3 推理系列产品:数据类型支持int8uint8int32bfloat16
    • div_modeFalse时,数据格式支持NDND。支持1维或多维(1维时,对应轴的大小需要与input中第axis维相等或等于1;多维时,scales的shape需要与input维度相等,除axis指定的维度,其他维度为1,axis指定的维度必须和input对应的维度相等)。zero_points的shape和dtype需要和scales一致。

      • Atlas 推理系列产品:数据类型支持floatfloat16
      • Atlas A2 训练系列产品/Atlas A2 推理系列产品、Atlas A3 训练系列产品/Atlas A3 推理系列产品:数据类型支持floatfloat16bfloat16。当input的数据格式为NZNZ时,值为空。
  • dtype (int):必选参数,指定输出参数的类型。

    • div_modeTrue时,

      • Atlas 推理系列产品:类型支持qint8quint8int32
      • Atlas A2 训练系列产品/Atlas A2 推理系列产品、Atlas A3 训练系列产品/Atlas A3 推理系列产品:类型支持qint8quint8int32
    • div_modeFalse时,类型支持qint8quint4x2。如果dtypequint4x2时,输出tensor类型为int32,由8个int4拼接。

  • axis (int):可选参数,量化的element-wise轴,其他的轴做broadcast,默认值为1

    div_modeFalse时,axis取值范围是[-2, +∞)且指定的轴不能超过输入input的维度数。如果axis为-2,代表量化的element-wise轴是输入input的倒数第二根轴;如果axis大于-2,量化的element-wise轴是输入的最后一根轴。

  • div_mode (bool):可选参数,表示计算scales模式,对应公式中的div_mode。当div_modeTrue时,表示用除法计算scalesdiv_modeFalse时,表示用乘法计算scales,默认值为True

返回值说明

Tensor

对应公式中的result。数据类型由参数dtype指定,如果参数dtypequint4x2,输出的dtypeint32,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,dtypequint4x2
  • div_modeFalse时:
    • 支持Atlas A2 训练系列产品/Atlas A2 推理系列产品。
    • dtypequint4x2或者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)