PyTorch torch.quantized_batch_norm Function


Pytorch torch 参考手册Pytorch torch Reference Manual

torch.quantized_batch_normIs a function in PyTorch used to perform batch normalization on quantized tensors. This function is very useful during quantized model inference, allowing batch normalization operations while maintaining the advantages of quantization.

Function Definition

torch.quantized_batch_norm(input, weight, bias, mean, var, eps, output_scale, output_zero_point)

Parameter Description

  • input: input quantized tensor
  • weight: batch normalization scale parameter
  • bias: batch normalization bias parameter
  • mean: batch normalization mean
  • var: batch normalization variance
  • eps: small constant to prevent division by zero
  • output_scale: quantization scale of the output tensor
  • output_zero_point: quantization zero point of the output tensor

Usage Example

Example

import torch

# Create quantized input tensor
input = torch.quantize_per_tensor(torch.randn(1, 3, 4, 4), scale=0.1, zero_point=0, dtype=torch.quint8)

# Batch normalization parameters
weight = torch.ones(3)
bias = torch.zeros(3)
mean = torch.ones(3) * 0.5
var = torch.ones(3) * 0.2

# Perform quantized batch normalization
output = torch.quantized_batch_norm(
    input, weight, bias, mean, var,
    eps=1e-5, output_scale=0.1, output_zero_point=0
)

print("Output shape:", output.shape)
print("Output type:", output.dtype)

The output result is:

输出形状: torch.Size([1, 3, 4, 4])
输出类型: torch.quint8

Pytorch torch 参考手册Pytorch torch Reference Manual

Other Extensions