РЕПЕТИТОР математика физика информатика
Для школьников и студентов. Подтягивание пробелов. ЦЭ, ЦТ, ОГЭ, ЕГЭ.
Идет набор на ЛЕТО. Жмите для подробностей:)
503 of 769 menu

Класс ParameterDict

Класс ParameterDict из модуля torch.nn представляет собой словарь, предназначенный для хранения параметров модели в виде объектов Parameter. Он наследуется от torch.nn.Module и позволяет обращаться к параметрам по строковым ключам, как к обычным атрибутам. Это особенно удобно, когда количество параметров динамическое или их имена формируются во время выполнения программы.

Синтаксис

torch.nn.ParameterDict(parameters=None)

Основные возможности

Класс ParameterDict поддерживает все операции обычного словаря Python: добавление, удаление, изменение элементов, а также итерацию по ключам и значениям. При этом все содержащиеся в нём параметры автоматически регистрируются как параметры модуля и будут учитываться при вызове parameters.

Пример

Давайте создадим пустой словарь параметров и добавим в него два параметра с разными именами:

import torch from torch.nn import ParameterDict, Parameter params = ParameterDict() params['weight'] = Parameter(torch.randn(3, 4)) params['bias'] = Parameter(torch.zeros(4)) for name, param in params.items(): print(f"{name}: {param.shape}")

Результат выполнения кода:

weight: torch.Size([3, 4]) bias: torch.Size([4])

Пример

Инициализируем словарь параметров сразу при создании, передав обычный словарь:

import torch from torch.nn import ParameterDict, Parameter torch.manual_seed(0) params = ParameterDict({ 'w1': Parameter(torch.ones(2, 3)), 'w2': Parameter(torch.randn(3, 2)), }) for name, param in params.items(): print(f"{name}:\\n{param}\\n")

Результат выполнения кода:

w1: tensor([[1., 1., 1.], [1., 1., 1.]]) w2: tensor([[ 1.5410, -0.2934], [-2.1788, 0.5684], [-1.0845, -1.3986]])

Пример

Продемонстрируем, как ParameterDict автоматически регистрирует параметры. Доступ к ним можно получить через метод parameters:

import torch from torch.nn import ParameterDict, Parameter params = ParameterDict() params['weight'] = Parameter(torch.ones(2, 2)) params['bias'] = Parameter(torch.zeros(2)) for param in params.parameters(): print(param)

Результат выполнения кода:

tensor([[1., 1.], [1., 1.]]) tensor([0., 0.])

Пример

Словарь параметров можно использовать внутри пользовательского модуля, чтобы динамически управлять параметрами:

import torch from torch.nn import Module, ParameterDict, Parameter class DynamicLinear(Module): def __init__(self, in_features, out_features): super().__init__() self.params = ParameterDict({ 'weight': Parameter(torch.randn(out_features, in_features)), 'bias': Parameter(torch.zeros(out_features)), }) def forward(self, x): return x @ self.params['weight'].T + self.params['bias'] torch.manual_seed(0) model = DynamicLinear(3, 2) x = torch.randn(1, 3) res = model(x) print(res)

Результат выполнения кода:

tensor([[-0.3252, 3.0241]])

Смотрите также

  • класс ParameterList,
    который хранит параметры в виде списка
  • класс ModuleDict,
    который хранит модули в виде словаря
  • класс UninitializedParameter,
    который представляет неинициализированный параметр
  • класс Buffer,
    который представляет буфер, не являющийся параметром
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить