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

Словарь весов модуля в PyTorch

Внутри модуля обучаемые числа имеют имена: у линейного слоя это матрица весов и вектор смещения. Метод state_dict собирает их в обычный словарь Python. Создадим простую сеть и выведем ключи:

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) net = TinyNet() weights = net.state_dict() print(list(weights.keys())) # выведет ['layer.weight', 'layer.bias']

Каждый ключ указывает на параметр внутри модели, а значение - тензор с числами. Такой словарь удобно передавать дальше, не таская за собой весь объект сети:

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) net = TinyNet() weights = net.state_dict() print(weights["layer.weight"].shape) # выведет torch.Size([3, 2])

Соберите модуль с линейным слоем на 2 входа и 4 выхода. Получите словарь его параметров и выведите список имён ключей.

Опишите блок с линейным преобразованием 3 на 1. Выведите форму тензора, который лежит по ключу с именем смещения слоя.

Создайте сеть с одним линейным слоем 1 на 2. Снимите с неё словарь параметров и выведите, сколько ключей в нём находится.

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