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

Загрузка весов в PyTorch

Если нужно вернуть сохранённые числа, файл читают функцией torch.load. Из файла берут только веса, а не произвольный код. Затем метод load_state_dict у нового модуля той же архитектуры подставляет тензоры на место параметров:

import torch import torch.nn as nn class TinyNet(nn.Module): def __init__(self): super().__init__() self.layer = nn.Linear(2, 3) def forward(self, x): return self.layer(x) torch.manual_seed(0) net = TinyNet() path = "weights.pt" torch.save(net.state_dict(), path) fresh = TinyNet() before = fresh.layer.weight[0, 0].item() stored = torch.load(path, weights_only=True) fresh.load_state_dict(stored) after = fresh.layer.weight[0, 0].item() print(before, after) # выведет 0.1871093362569809 -0.005293981172144413

После подстановки веса совпадают с исходной моделью. Сравним матрицу линейного слоя у сохранённой и загруженной сети:

import torch import torch.nn as nn class TinyNet(nn.Module): def __init__(self): super().__init__() self.layer = nn.Linear(2, 3) def forward(self, x): return self.layer(x) torch.manual_seed(0) net = TinyNet() path = "weights.pt" torch.save(net.state_dict(), path) fresh = TinyNet() stored = torch.load(path, weights_only=True) fresh.load_state_dict(stored) same = torch.allclose( fresh.layer.weight, net.layer.weight ) print(same) # выведет True

Создайте модуль с линейным слоем 1 на 1, запишите его параметры в one.pt. Соберите второй такой же модуль, прочитайте файл и подставьте данные в параметры. Выведите один элемент матрицы весов до подстановки и после неё.

Опишите сеть с линейным преобразованием 2 на 2, сохраните веса в pair.pt. На новом объекте той же формы восстановите параметры из файла и выведите, совпадают ли матрицы весов.

Соберите блок 3 на 1, положите словарь параметров в bias.pt. Создайте ещё один блок той же схемы, загрузите данные с диска и выведите форму вектора смещения после восстановления.

← →
↑
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить