Spec-Zone.ru › PyTorch 2.14

torch.set_default_dtype

torch.set_default_dtype(d, /) [исходный код]

Устанавливает тип данных с плавающей точкой по умолчанию в значение d. В качестве входных данных поддерживаются типы данных с плавающей точкой. Для других типов данных torch вызовет исключение.

При инициализации PyTorch типом данных с плавающей точкой по умолчанию является torch.float32. Функция set_default_dtype(torch.float64) предназначена для упрощения вывода типов, подобного NumPy. Тип данных с плавающей точкой по умолчанию используется для:

  1. Неявного определения комплексного типа данных по умолчанию. Если тип данных с плавающей точкой по умолчанию — float16, комплексным типом данных по умолчанию будет complex32. Для float32 комплексным типом данных по умолчанию будет complex64. Для float64 — complex128. Для bfloat16 будет вызвано исключение, так как для bfloat16 не существует соответствующего комплексного типа.
  2. Вывода типа данных для тензоров, созданных с использованием чисел с плавающей точкой или комплексных чисел Python. См. примеры ниже.
  3. Определения результата продвижения типов между логическими и целочисленными тензорами, а также числами с плавающей точкой и комплексными числами 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

Spec-Zone.ru

Настройки Оффлайн Что нового Помощь О нас
Spec-Zone .ru
спецификации, руководства, описания, API