PyTorch torch.quantized_batch_norm Function
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 tensorweight: batch normalization scale parameterbias: batch normalization bias parametermean: batch normalization meanvar: batch normalization varianceeps: small constant to prevent division by zerooutput_scale: quantization scale of the output tensoroutput_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)
# 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
Other Extensions