Метод extra_repr
Метод extra_repr класса Module используется для
формирования дополнительной строковой информации о модуле.
Этот метод вызывается автоматически при выводе модуля
через print или при преобразовании в строку.
Метод не принимает параметров и должен возвращать строку.
Синтаксис
def extra_repr(self) -> str:
return "additional info"
Пример
Давайте создадим простой линейный слой и посмотрим его стандартное строковое представление:
import torch
layer = torch.nn.Linear(10, 5)
print(layer)
Результат выполнения кода:
Linear(in_features=10, out_features=5, bias=True)
Пример
Создадим собственный модуль, переопределяющий метод
extra_repr для добавления пользовательской
информации:
import torch
class MyModule(torch.nn.Module):
def __init__(self, input_size, output_size):
super().__init__()
self.input_size = input_size
self.output_size = output_size
self.linear = torch.nn.Linear(input_size, output_size)
def extra_repr(self):
return f"input_size={self.input_size}, output_size={self.output_size}"
module = MyModule(10, 5)
print(module)
Результат выполнения кода:
MyModule(
(linear): Linear(in_features=10, out_features=5, bias=True)
)input_size=10, output_size=5
Пример
Добавим в extra_repr отображение состояния
обучения модуля:
import torch
class CustomLayer(torch.nn.Module):
def __init__(self, dim, dropout=0.0):
super().__init__()
self.dim = dim
self.dropout = dropout
self.weight = torch.nn.Parameter(torch.randn(dim, dim))
def extra_repr(self):
training_status = "train" if self.training else "eval"
return f"dim={self.dim}, dropout={self.dropout}, status={training_status}"
layer = CustomLayer(3, 0.5)
print(layer)
layer.eval()
print(layer)
Результат выполнения кода:
CustomLayer(
(weight): Parameter containing:
[torch.FloatTensor of size 3x3]
)dim=3, dropout=0.5, status=train
CustomLayer(
(weight): Parameter containing:
[torch.FloatTensor of size 3x3]
)dim=3, dropout=0.5, status=eval
Пример
Используем extra_repr для отображения
гиперпараметров сложного модуля:
import torch
class ComplexModule(torch.nn.Module):
def __init__(self, in_dim, hidden_dim, out_dim, activation="relu"):
super().__init__()
self.in_dim = in_dim
self.hidden_dim = hidden_dim
self.out_dim = out_dim
self.activation = activation
self.fc1 = torch.nn.Linear(in_dim, hidden_dim)
self.fc2 = torch.nn.Linear(hidden_dim, out_dim)
def extra_repr(self):
return (f"in_dim={self.in_dim}, hidden_dim={self.hidden_dim}, "
f"out_dim={self.out_dim}, activation={self.activation}")
module = ComplexModule(20, 50, 10, "gelu")
print(module)
Результат выполнения кода:
ComplexModule(
(fc1): Linear(in_features=20, out_features=50, bias=True)
(fc2): Linear(in_features=50, out_features=10, bias=True)
)in_dim=20, hidden_dim=50, out_dim=10, activation=gelu