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

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