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

Метод 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,
    который возвращает итератор по всем буферам модуля
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить