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

Метод buffers

Метод buffers применяется к экземпляру класса Module и возвращает генератор, который итерируется по всем буферам, зарегистрированным в модуле. Буферы - это тензоры, которые сохраняются внутри модели, но не обновляются оптимизатором (в отличие от параметров). Типичные примеры буферов - тензоры с running mean и running variance в слоях батч-нормализации. Метод не принимает аргументов и возвращает итератор по тензорам.

Синтаксис

module.buffers()

Пример

Давайте создадим простой модуль с одним буфером и выведем все его буферы с помощью метода buffers:

import torch import torch.nn as nn class MyModule(nn.Module): def __init__(self): super().__init__() self.register_buffer('my_buffer', torch.tensor([1.0, 2.0, 3.0])) model = MyModule() for buf in model.buffers(): print(buf)

Результат выполнения кода:

tensor([1., 2., 3.])

Пример

Теперь рассмотрим модуль, содержащий как параметры, так и буферы. Метод buffers возвращает только буферы, не затрагивая обучаемые параметры:

import torch import torch.nn as nn class MyModule(nn.Module): def __init__(self): super().__init__() self.param = nn.Parameter(torch.tensor([1.0, 2.0])) self.register_buffer('running_mean', torch.tensor([0.5, 0.5])) self.register_buffer('running_var', torch.tensor([1.0, 1.0])) model = MyModule() print("Buffers:") for buf in model.buffers(): print(buf)

Результат выполнения кода:

Buffers: tensor([0.5000, 0.5000]) tensor([1., 1.])

Пример

Метод buffers также работает рекурсивно для вложенных модулей, перебирая все буферы во всех дочерних подмодулях:

import torch import torch.nn as nn class SubModule(nn.Module): def __init__(self): super().__init__() self.register_buffer('sub_buf', torch.tensor([9.0, 8.0])) class MainModule(nn.Module): def __init__(self): super().__init__() self.sub = SubModule() self.register_buffer('main_buf', torch.tensor([1.0, 2.0])) model = MainModule() for buf in model.buffers(): print(buf)

Результат выполнения кода:

tensor([1., 2.]) tensor([9., 8.])

Смотрите также

  • метод named_buffers,
    который возвращает имена и буферы модуля
  • метод parameters,
    который возвращает обучаемые параметры модуля
  • метод state_dict,
    который возвращает состояние модуля, включая буферы
  • метод register_buffer,
    который регистрирует новый буфер в модуле
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить