torch.set_default_dtype
-
torch.set_default_dtype(d)[source] -
Устанавливает тип по умолчанию для чисел с плавающей запятой в
d. Поддерживаются типы torch.float32 и torch.float64. Другие типы могут быть приняты без замечаний, но не поддерживаются и вряд ли будут работать как ожидается.При инициализации PyTorch тип по умолчанию для чисел с плавающей запятой — torch.float32, и цель set_default_dtype(torch.float64) — облегчить инференс типа, подобный NumPy. Тип по умолчанию для чисел с плавающей запятой используется для:
- Неявного определения типа по умолчанию для комплексных чисел. Если тип по умолчанию для чисел с плавающей запятой — float32, то тип по умолчанию для комплексных чисел — complex64, а если тип по умолчанию для чисел с плавающей запятой — float64, то тип по умолчанию для комплексных чисел — complex128.
- Вывода типа для тензоров, созданных с использованием Python-чисел с плавающей запятой или комплексных Python-чисел. См. примеры ниже.
- Определения результата повышения типа между тензорами bool и целыми числами и Python-числами с плавающей запятой и комплексными Python-числами.
- Parameters:
-
d (
torch.dtype) — тип чисел с плавающей запятой, который должен стать типом по умолчанию. Это torch.float32 или torch.float64.
Пример
>>> # 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
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/1.13/generated/torch.set_default_dtype.html