Класс Identity
Класс Identity из модуля torch.nn представляет собой слой, который выполняет тождественное преобразование. Данный слой принимает входной тензор и возвращает его без каких-либо изменений. Он полезен в ситуациях, когда необходимо сохранить архитектуру сети, но при этом требуется пропустить какой-либо этап обработки, либо когда нужно оставить запасной путь в графе вычислений.
Класс не принимает никаких параметров при инициализации. Он может использоваться как самостоятельный слой в последовательной модели, так и как часть более сложных архитектур, например, в блоках типа ResNet.
Синтаксис
torch.nn.Identity(*args, **kwargs)
Объект класса создаётся без обязательных аргументов. При вызове слоя он принимает произвольное количество аргументов, но возвращает только первый переданный тензор.
Пример
Давайте создадим слой тождественного преобразования и применим его к тензору:
import torch
import torch.nn as nn
# Создаём слой Identity
identity = nn.Identity()
# Создаём входной тензор
t = torch.tensor([1, 2, 3, 4, 5])
# Применяем слой
res = identity(t)
print(res)
Результат выполнения кода:
tensor([1, 2, 3, 4, 5])
Пример
Рассмотрим использование Identity внутри последовательной модели Sequential. Это может быть полезно, например, для пропуска этапа активации:
import torch
import torch.nn as nn
# Создаём последовательную модель
model = nn.Sequential(
nn.Linear(10, 5),
nn.Identity(), # Пропускаем активацию
nn.Linear(5, 2)
)
# Создаём входной тензор
t = torch.randn(3, 10)
# Применяем модель
res = model(t)
print(res.shape)
Результат выполнения кода:
torch.Size([3, 2])
Пример
Слой Identity может принимать несколько аргументов, но обрабатывает только первый. Это полезно, когда слой используется в функциях, ожидающих несколько входных тензоров:
import torch
import torch.nn as nn
identity = nn.Identity()
t1 = torch.tensor([1, 2, 3])
t2 = torch.tensor([4, 5, 6])
t3 = torch.tensor([7, 8, 9])
# Передаём три тензора, возвращается только первый
res = identity(t1, t2, t3)
print(res)
Результат выполнения кода:
tensor([1, 2, 3])