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

Метод state_dict

Метод state_dict класса Module возвращает словарь Python, содержащий все обучаемые параметры (веса и смещения) и постоянные буферы модели. Ключами словаря выступают строковые имена параметров, значениями - тензоры. Метод не принимает параметров.

Синтаксис

model.state_dict()

Пример

Создадим простую линейную модель и посмотрим её параметры:

import torch import torch.nn as nn model = nn.Linear(3, 2) state = model.state_dict() for key, value in state.items(): print(key, value.shape)

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

weight torch.Size([2, 3]) bias torch.Size([2])

Как видно из результата, модель содержит два параметра - матрицу весов размером 2x3 и вектор смещения размером 2.

Пример

Создадим более сложную модель с несколькими слоями:

import torch import torch.nn as nn class SimpleNN(nn.Module): def __init__(self): super().__init__() self.fc1 = nn.Linear(10, 5) self.fc2 = nn.Linear(5, 1) self.relu = nn.ReLU() def forward(self, x): x = self.relu(self.fc1(x)) return self.fc2(x) model = SimpleNN() state = model.state_dict() for key, value in state.items(): print(key, value.shape)

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

fc1.weight torch.Size([5, 10]) fc1.bias torch.Size([5]) fc2.weight torch.Size([1, 5]) fc2.bias torch.Size([1])

Пример

Сохраним параметры модели в файл и восстановим их в другой модели:

import torch import torch.nn as nn model1 = nn.Linear(5, 3) model2 = nn.Linear(5, 3) torch.save(model1.state_dict(), 'model.pt') model2.load_state_dict(torch.load('model.pt')) print(torch.equal(model1.weight, model2.weight))

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

True

В этом примере мы сохранили состояние первой модели в файл, а затем загрузили его во вторую модель. Проверка показала, что параметры полностью совпадают.

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

  • метод load_state_dict,
    который загружает параметры модели из словаря
  • метод parameters,
    который возвращает итератор по параметрам модели
  • метод named_parameters,
    который возвращает итератор с именами и параметрами модели
  • метод buffers,
    который возвращает итератор по постоянным буферам модели
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить