Метод get_buffer
Метод get_buffer класса Module возвращает буфер модуля
по указанному строковому пути. Буферы - это тензоры,
которые хранятся в состоянии модели, но не обновляются
во время градиентного спуска, в отличие от параметров.
Метод принимает один обязательный аргумент - строку
с путём к буферу, где имена подмодулей разделяются
точками. Если буфер с указанным путём не найден,
метод выбрасывает исключение AttributeError.
Синтаксис
module.get_buffer(target)
Пример
Создадим простой модуль с зарегистрированным буфером
и получим его с помощью get_buffer:
import torch
from torch import nn
class MyModule(nn.Module):
def __init__(self):
super().__init__()
self.register_buffer('running_mean', torch.zeros(5))
model = MyModule()
buffer = model.get_buffer('running_mean')
print(buffer)
Результат выполнения кода:
tensor([0., 0., 0., 0., 0.])
Пример
Получим буфер из вложенного подмодуля, используя путь с точками:
import torch
from torch import nn
class SubModule(nn.Module):
def __init__(self):
super().__init__()
self.register_buffer('weight', torch.ones(3, 3))
class MainModule(nn.Module):
def __init__(self):
super().__init__()
self.sub = SubModule()
model = MainModule()
buffer = model.get_buffer('sub.weight')
print(buffer)
Результат выполнения кода:
tensor([
[1., 1., 1.],
[1., 1., 1.],
[1., 1., 1.],
])
Пример
Попробуем получить несуществующий буфер, что вызовет ошибку:
import torch
from torch import nn
class MyModule(nn.Module):
def __init__(self):
super().__init__()
self.register_buffer('running_mean', torch.zeros(5))
model = MyModule()
try:
buffer = model.get_buffer('nonexistent')
except AttributeError as e:
print(e)
Результат выполнения кода:
"type object 'MyModule' has no attribute or buffer 'nonexistent'"
Смотрите также
-
метод
buffers,
который возвращает итератор по всем буферам модуля -
метод
named_buffers,
который возвращает итератор по именам и буферам модуля -
метод
register_buffer,
который регистрирует буфер в модуле -
метод
get_parameter,
который возвращает параметр модуля по указанному пути