PyTorch torch.set_default_dtype Function
Pytorch torch Reference Manual
torch.set_default_dtypeIt is a function in PyTorch used to set the default floating-point data type.
Function Definition
torch.set_default_dtype(d)
Usage Example
Example
import torch
# Set the default dtype to float64
torch.set_default_dtype(torch.float64)
# Create a tensor, using float64 by default
x = torch.tensor([1.0, 2.0, 3.0])
print("Default dtype:", x.dtype)
# Restore default
torch.set_default_dtype(torch.float32)
# Set the default dtype to float64
torch.set_default_dtype(torch.float64)
# Create a tensor, using float64 by default
x = torch.tensor([1.0, 2.0, 3.0])
print("Default dtype:", x.dtype)
# Restore default
torch.set_default_dtype(torch.float32)
Other Extensions