Метод 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,
который регистрирует новый буфер в модуле