PyTorch torch.set_float32_matmul_precision Function


Pytorch torch 参考手册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")

The output is:

使用中等精度进行矩阵乘法
结果形状: torch.Size([100, 100])

Pytorch torch 参考手册Pytorch torch Reference Manual

Other Extensions