Метод register_buffer
Метод register_buffer класса Module предназначен для регистрации тензора в составе модуля в качестве буфера.
Буферы - это тензоры, которые не являются обучаемыми параметрами модели (не обновляются во время обучения через градиенты),
но при этом должны сохраняться в состоянии модели (state_dict) и перемещаться вместе с модулем при вызове методов to, cuda и других.
Метод принимает два обязательных аргумента: имя буфера и тензор, а также необязательный флаг, определяющий, является ли буфер постоянным.
Синтаксис
module.register_buffer(name, tensor, persistent=True)
Пример
Давайте создадим простой модуль и зарегистрируем в нем буфер с постоянным значением:
import torch
import torch.nn as nn
class MyModule(nn.Module):
def __init__(self):
super().__init__()
self.register_buffer('my_buffer', torch.tensor([1, 2, 3, 4, 5]))
model = MyModule()
print(model.my_buffer)
Результат выполнения кода:
tensor([1, 2, 3, 4, 5])
Пример
Буферы могут использоваться для хранения средних значений или статистик, которые не должны обновляться через оптимизатор. Рассмотрим модуль, который хранит скользящее среднее:
import torch
import torch.nn as nn
class RunningAverageModule(nn.Module):
def __init__(self):
super().__init__()
self.register_buffer('running_avg', torch.zeros(1))
self.register_buffer('count', torch.zeros(1))
def update(self, value):
self.count += 1
self.running_avg = self.running_avg + (value - self.running_avg) / self.count
model = RunningAverageModule()
model.update(torch.tensor([5.0]))
model.update(torch.tensor([7.0]))
print(model.running_avg)
Результат выполнения кода:
tensor([6.])
Пример
Буферы также можно использовать для хранения весов, которые не должны обучаться, например, маски или фиксированные сверточные ядра. В этом примере мы создадим модуль, который применяет фиксированное ядро к входным данным:
import torch
import torch.nn as nn
import torch.nn.functional as F
class FixedKernelConv(nn.Module):
def __init__(self):
super().__init__()
kernel = torch.tensor([[[[1.0, 0.0, -1.0]]]])
self.register_buffer('kernel', kernel)
def forward(self, x):
return F.conv2d(x, self.kernel)
model = FixedKernelConv()
x = torch.randn(1, 1, 3, 3)
res = model(x)
print(res.shape)
Результат выполнения кода:
torch.Size([1, 1, 3, 1])
Пример
Буферы сохраняются в state_dict модуля, что позволяет сохранять их вместе с моделью.
Рассмотрим пример сохранения и загрузки буфера:
import torch
import torch.nn as nn
class MyModule(nn.Module):
def __init__(self):
super().__init__()
self.register_buffer('buffer', torch.tensor([1, 2, 3]))
model = MyModule()
print('Before save:', model.buffer)
state = model.state_dict()
model.buffer = torch.tensor([4, 5, 6])
print('After modification:', model.buffer)
model.load_state_dict(state)
print('After load:', model.buffer)
Результат выполнения кода:
Before save: tensor([1, 2, 3])
After modification: tensor([4, 5, 6])
After load: tensor([1, 2, 3])
Смотрите также
-
метод
state_dict,
который возвращает словарь всех параметров и буферов модуля -
метод
load_state_dict,
который загружает параметры и буферы из словаря -
метод
register_parameter,
который регистрирует обучаемый параметр модуля -
метод
buffers,
который возвращает итератор по всем буферам модуля