Метод float
Метод float класса Module преобразует все параметры модели
(веса и смещения) и буферы в тип данных float32. Метод не принимает
параметров и возвращает сам объект модели, что позволяет использовать
цепочки вызовов. Это особенно полезно при работе с моделями,
изначально созданными в другом типе данных, например float16 или
float64, или при необходимости обеспечить единообразие вычислений.
Синтаксис
model.float()
Пример
Давайте создадим простую линейную модель и приведём её параметры к типу float:
import torch
import torch.nn as nn
model = nn.Linear(10, 5)
print(model.weight.dtype)
model = model.float()
print(model.weight.dtype)
Результат выполнения кода:
torch.float32
torch.float32
Пример
Создадим модель с параметрами типа float64 и преобразуем её в float32:
import torch
import torch.nn as nn
model = nn.Linear(10, 5)
model = model.double()
print(model.weight.dtype)
model = model.float()
print(model.weight.dtype)
Результат выполнения кода:
torch.float64
torch.float32
Пример
Метод float можно использовать в цепочке с другими методами,
например перед перемещением модели на устройство:
import torch
import torch.nn as nn
model = nn.Sequential(
nn.Linear(10, 20),
nn.ReLU(),
nn.Linear(20, 5)
)
model = model.float().to('cpu')
print(model[0].weight.dtype)
Результат выполнения кода:
torch.float32