Тип float32
Тип данных torch.float32 (также доступен как torch.float)
представляет 32-битное число с плавающей точкой. Этот тип используется
по умолчанию при создании тензоров из чисел с плавающей точкой.
Он обеспечивает хороший баланс между точностью вычислений и
потреблением памяти, что делает его наиболее популярным типом
для большинства задач глубокого обучения.
Создание тензора с типом float32
Укажем тип данных float32 при создании тензора:
import torch
t = torch.tensor([1, 2, 3, 4, 5], dtype=torch.float32)
print(t)
print(t.dtype)
Результат выполнения кода:
tensor([1., 2., 3., 4., 5.])
torch.float32
Обратите внимание, что числа в выводе отображаются с точкой, что указывает на тип с плавающей точкой.
Синтаксис
Тип данных float32 используется в параметре
dtype функций создания тензоров:
torch.tensor(data, dtype=torch.float32)
torch.zeros(size, dtype=torch.float32)
torch.ones(size, dtype=torch.float32)
Также тип можно указать при преобразовании существующего тензора
с помощью метода to:
t = torch.tensor([1, 2, 3])
t = t.to(torch.float32)
Или с помощью метода float:
t = torch.tensor([1, 2, 3])
t = t.float()
Для проверки типа данных тензора используйте атрибут
dtype:
print(t.dtype)
Результат выполнения кода:
torch.float32
Пример с тензором вещественных чисел
При создании тензора из вещественных чисел тип float32
используется автоматически:
import torch
t = torch.tensor([0.1, 0.2, 0.3, 0.4, 0.5])
print(t)
print(t.dtype)
Результат выполнения кода:
tensor([0.1000, 0.2000, 0.3000, 0.4000, 0.5000])
torch.float32
Преобразование целых чисел в float32
Преобразуем тензор целых чисел в тип float32
с помощью метода float:
import torch
t_int = torch.tensor([1, 2, 3, 4, 5], dtype=torch.int32)
t_float = t_int.float()
print(t_int)
print(t_float)
Результат выполнения кода:
tensor([1, 2, 3, 4, 5], dtype=torch.int32)
tensor([1., 2., 3., 4., 5.])
Проверка совместимости типов
Проверим, можно ли безопасно преобразовать тип int32
в float32 с помощью функции can_cast:
import torch
res = torch.can_cast(torch.int32, torch.float32)
print(res)
Результат выполнения кода:
True
Смотрите также
-
тип
float16,
который использует 16 бит для хранения чисел с плавающей точкой -
тип
float64,
который использует 64 бита для повышенной точности вычислений -
тип
bfloat16,
который представляет 16-битный формат с расширенным диапазоном -
функцию
get_default_dtype,
которая возвращает текущий тип данных по умолчанию