PyTorch torch.set_float32_matmul_precision Function
Pytorch torch Reference Manual
torch.set_float32_matmul_precisionis a function in PyTorch used to set the precision of float32 matrix multiplication. You can choose to use lower precision to improve performance, or use higher precision to improve accuracy.
Function Definition
torch.set_float32_matmul_precision(precision)
Parameter Description
precision: Precision level, optional values:"highest": Highest precision (default)"high": High precision"medium": Medium precision (uses TensorFloat-32)
Usage Example
Example
import torch
# Set to medium precision (use TensorFloat-32 for acceleration)
torch.set_float32_matmul_precision("medium")
# Create matrices for testing
a = torch.randn(100, 100)
b = torch.randn(100, 100)
# Matrix multiplication
c = torch.matmul(a, b)
print("Matrix multiplication using medium precision")
print("Result shape:", c.shape)
# Restore to highest precision
torch.set_float32_matmul_precision("highest")
# Set to medium precision (use TensorFloat-32 for acceleration)
torch.set_float32_matmul_precision("medium")
# Create matrices for testing
a = torch.randn(100, 100)
b = torch.randn(100, 100)
# Matrix multiplication
c = torch.matmul(a, b)
print("Matrix multiplication using medium precision")
print("Result shape:", c.shape)
# Restore to highest precision
torch.set_float32_matmul_precision("highest")
The output is:
使用中等精度进行矩阵乘法 结果形状: torch.Size([100, 100])
Other Extensions