torch.set_default_dtype
-
torch.set_default_dtype(d, /)[исходный код] -
Устанавливает тип данных с плавающей точкой по умолчанию в значение
d. В качестве входных данных поддерживаются типы данных с плавающей точкой. Для других типов данных torch вызовет исключение.При инициализации PyTorch типом данных с плавающей точкой по умолчанию является torch.float32. Функция set_default_dtype(torch.float64) предназначена для упрощения вывода типов, подобного NumPy. Тип данных с плавающей точкой по умолчанию используется для:
- Неявного определения комплексного типа данных по умолчанию. Если тип данных с плавающей точкой по умолчанию — float16, комплексным типом данных по умолчанию будет complex32. Для float32 комплексным типом данных по умолчанию будет complex64. Для float64 — complex128. Для bfloat16 будет вызвано исключение, так как для bfloat16 не существует соответствующего комплексного типа.
- Вывода типа данных для тензоров, созданных с использованием чисел с плавающей точкой или комплексных чисел Python. См. примеры ниже.
- Определения результата продвижения типов между логическими и целочисленными тензорами, а также числами с плавающей точкой и комплексными числами Python.
- Параметры:
-
d (
torch.dtype) – тип данных с плавающей точкой, который нужно установить по умолчанию.
Пример
>>> # initial default for floating point is torch.float32 >>> # Python floats are interpreted as float32 >>> torch.tensor([1.2, 3]).dtype torch.float32 >>> # initial default for floating point is torch.complex64 >>> # Complex Python numbers are interpreted as complex64 >>> torch.tensor([1.2, 3j]).dtype torch.complex64
>>> torch.set_default_dtype(torch.float64) >>> # Python floats are now interpreted as float64 >>> torch.tensor([1.2, 3]).dtype # a new floating point tensor torch.float64 >>> # Complex Python numbers are now interpreted as complex128 >>> torch.tensor([1.2, 3j]).dtype # a new complex tensor torch.complex128
>>> torch.set_default_dtype(torch.float16) >>> # Python floats are now interpreted as float16 >>> torch.tensor([1.2, 3]).dtype # a new floating point tensor torch.float16 >>> # Complex Python numbers are now interpreted as complex128 >>> torch.tensor([1.2, 3j]).dtype # a new complex tensor torch.complex32
© 2026, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://docs.pytorch.org/docs/2.14/generated/torch.set_default_dtype.html