Метод double
Метод double класса Module преобразует все параметры модели
и постоянные буферы в тип данных с двойной точностью torch.float64.
Этот метод полезен, когда требуется повысить точность вычислений
или работать с данными в формате double. Метод изменяет модель
на месте и возвращает саму модель для удобства вызова цепочкой.
Синтаксис
model.double()
Пример
Давайте создадим простую линейную модель и преобразуем её параметры в тип double:
import torch
import torch.nn as nn
model = nn.Linear(5, 3)
print(model.weight.dtype)
model = model.double()
print(model.weight.dtype)
Результат выполнения кода:
torch.float32
torch.float64
Как видите, после вызова метода double тип данных параметров
изменился с float32 на float64.
Пример
Давайте применим метод double к последовательной модели,
состоящей из нескольких слоёв:
import torch
import torch.nn as nn
model = nn.Sequential(
nn.Linear(10, 20),
nn.ReLU(),
nn.Linear(20, 5)
)
print(model[0].weight.dtype)
model.double()
print(model[0].weight.dtype)
print(model[2].weight.dtype)
Результат выполнения кода:
torch.float32
torch.float64
torch.float64
Метод double рекурсивно применяется ко всем вложенным модулям,
поэтому все параметры модели становятся типа float64.
Пример
Давайте создадим модель с пользовательским модулем и убедимся,
что метод double работает с параметрами, созданными через Parameter:
import torch
import torch.nn as nn
class MyModule(nn.Module):
def __init__(self):
super().__init__()
self.weight = nn.Parameter(torch.randn(3, 4))
self.bias = nn.Parameter(torch.zeros(4))
def forward(self, x):
return x @ self.weight + self.bias
model = MyModule()
print(model.weight.dtype)
print(model.bias.dtype)
model.double()
print(model.weight.dtype)
print(model.bias.dtype)
Результат выполнения кода:
torch.float32
torch.float32
torch.float64
torch.float64
Метод double работает с любыми параметрами, зарегистрированными
в модели, включая созданные через Parameter.
Смотрите также
-
метод
float,
который преобразует параметры модели в тип float -
метод
half,
который преобразует параметры модели в тип half -
метод
type,
который преобразует параметры модели в указанный тип -
метод
parameters,
который возвращает все параметры модели