Метод get_parameter
Метод get_parameter класса Module возвращает параметр модели по его имени. Первым параметром метод принимает строку с полным именем параметра, включая путь к вложенным модулям. Если параметр с указанным именем не найден, метод выбрасывает исключение AttributeError. Данный метод удобен для динамического доступа к параметрам модели во время обучения и инференса.
Синтаксис
module.get_parameter(target)
Пример
Давайте создадим простую линейную модель и получим параметр весов с помощью метода get_parameter:
import torch
import torch.nn as nn
model = nn.Linear(5, 3)
weight = model.get_parameter('weight')
print(weight.shape)
Результат выполнения кода:
torch.Size([3, 5])
Пример
Давайте создадим модель с вложенными модулями и получим параметр из вложенного слоя:
import torch
import torch.nn as nn
class MyModel(nn.Module):
def __init__(self):
super().__init__()
self.layer1 = nn.Linear(10, 5)
self.layer2 = nn.Linear(5, 2)
model = MyModel()
weight = model.get_parameter('layer1.weight')
print(weight.shape)
Результат выполнения кода:
torch.Size([5, 10])
Пример
Давайте изменим значение параметра, полученного с помощью метода get_parameter:
import torch
import torch.nn as nn
torch.manual_seed(0)
model = nn.Linear(3, 2)
bias = model.get_parameter('bias')
print('Before:', bias)
bias.data.zero_()
print('After:', bias)
Результат выполнения кода:
Before: tensor([-0.2073, 0.4649], requires_grad=True)
After: tensor([0., 0.], requires_grad=True)
Смотрите также
-
метод
named_parameters,
который возвращает итератор по именам и параметрам модели -
метод
parameters,
который возвращает итератор по всем параметрам модели -
метод
register_parameter,
который регистрирует новый параметр в модели -
метод
state_dict,
который возвращает словарь всех параметров модели