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

Метод load_state_dict

Метод load_state_dict загружает состояние модели из словаря, полученного ранее с помощью метода state_dict. Первым параметром метод принимает словарь с весами, вторым параметром можно передать флаг строгой загрузки strict. Если strict=True (по умолчанию), то ключи загружаемого словаря должны точно совпадать с ключами текущей модели. Если некоторые ключи отсутствуют или лишние, будет вызвана ошибка.

Синтаксис

model.load_state_dict(state_dict, strict=True)

Пример

Давайте создадим простую модель, сохраним её состояние, а затем загрузим его в новую модель:

import torch from torch import nn class SimpleModel(nn.Module): def __init__(self): super().__init__() self.fc = nn.Linear(5, 3) def forward(self, x): return self.fc(x) model1 = SimpleModel() model2 = SimpleModel() state = model1.state_dict() model2.load_state_dict(state)

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

<All keys matched successfully>

Пример

Давайте загрузим состояние с проверкой ключей. Если структура модели не совпадает, будет ошибка:

import torch from torch import nn class ModelA(nn.Module): def __init__(self): super().__init__() self.fc = nn.Linear(5, 3) def forward(self, x): return self.fc(x) class ModelB(nn.Module): def __init__(self): super().__init__() self.fc1 = nn.Linear(5, 3) def forward(self, x): return self.fc1(x) model_a = ModelA() model_b = ModelB() state_a = model_a.state_dict() model_b.load_state_dict(state_a, strict=True)

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

RuntimeError: Error(s) in loading state_dict for ModelB: Missing key(s) in state_dict: "fc1.weight", "fc1.bias" Unexpected key(s) in state_dict: "fc.weight", "fc.bias"

Пример

Давайте загрузим состояние с отключённой строгой проверкой. Лишние или отсутствующие ключи будут проигнорированы:

import torch from torch import nn class ModelA(nn.Module): def __init__(self): super().__init__() self.fc = nn.Linear(5, 3) self.extra = nn.Linear(3, 1) def forward(self, x): return self.fc(x) class ModelB(nn.Module): def __init__(self): super().__init__() self.fc = nn.Linear(5, 3) def forward(self, x): return self.fc(x) model_a = ModelA() model_b = ModelB() state_a = model_a.state_dict() model_b.load_state_dict(state_a, strict=False)

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

<All keys matched successfully>

Пример

Давайте загрузим предварительно сохранённое состояние из файла. Обычно модель сохраняется с помощью метода save:

import torch from torch import nn class SimpleModel(nn.Module): def __init__(self): super().__init__() self.fc = nn.Linear(5, 3) def forward(self, x): return self.fc(x) model = SimpleModel() state_dict = torch.load('model.pt') model.load_state_dict(state_dict)

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

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