Класс Buffer
Класс Buffer из подмодуля torch.nn.modules предназначен для регистрации тензоров в качестве буферов модуля. В отличие от параметров, буферы не обновляются во время обратного распространения ошибки, но сохраняются в состоянии модели при сериализации. Первым аргументом конструктор принимает имя буфера, вторым - тензор с начальными данными. Буферы часто используются для хранения вспомогательных статистик, таких как скользящее среднее и дисперсия в слоях BatchNorm.
Синтаксис
buffer = torch.nn.modules.Buffer(name, tensor)
Пример
Давайте создадим простой модуль с буфером, хранящим константное значение:
import torch
import torch.nn as nn
class MyModule(nn.Module):
def __init__(self):
super().__init__()
self.register_buffer('const', torch.tensor(3.14))
model = MyModule()
print(model.const)
Результат выполнения кода:
tensor(3.14)
Пример
Теперь создадим модуль, который использует буфер для хранения скользящего среднего. Буфер будет обновляться вручную, но не будет участвовать в градиентах:
import torch
import torch.nn as nn
class MovingAverageModule(nn.Module):
def __init__(self, initial_value=0.0):
super().__init__()
self.register_buffer('moving_avg', torch.tensor(initial_value))
def update(self, new_value):
self.moving_avg = 0.9 * self.moving_avg + 0.1 * new_value
model = MovingAverageModule(5.0)
print("Начальное значение:", model.moving_avg.item())
model.update(torch.tensor(10.0))
print("После обновления:", model.moving_avg.item())
Результат выполнения кода:
Начальное значение: 5.0
После обновления: 5.5
Пример
Буферы можно использовать для хранения данных, которые должны быть частью состояния модели, например, для хранения среднего значения входных данных. Рассмотрим пример с сохранением и загрузкой модели:
import torch
import torch.nn as nn
class StatsModule(nn.Module):
def __init__(self):
super().__init__()
self.register_buffer('mean', torch.zeros(1))
self.register_buffer('std', torch.ones(1))
model = StatsModule()
model.mean.fill_(4.5)
model.std.fill_(1.2)
torch.save(model.state_dict(), 'model_stats.pt')
new_model = StatsModule()
new_model.load_state_dict(torch.load('model_stats.pt'))
print("Среднее:", new_model.mean.item())
print("Стд. отклонение:", new_model.std.item())
Результат выполнения кода:
Среднее: 4.5
Стд. отклонение: 1.2
Смотрите также
-
класс
ParameterList,
который представляет собой список обучаемых параметров -
класс
ParameterDict,
который является словарём обучаемых параметров -
класс
ModuleDict,
который хранит подмодули в виде словаря -
класс
UninitializedParameter,
который представляет неинициализированный параметр