Класс Hardswish
Класс Hardswish реализует функцию активации, которая является кусочно-линейной аппроксимацией функции swish. Функция применяется к каждому элементу тензора по формуле: x * min(6, max(0, x + 3)) / 6. Класс не имеет обучаемых параметров.
Синтаксис
torch.nn.Hardswish()
Пример
Давайте создадим слой Hardswish и применим его к тензору:
import torch
layer = torch.nn.Hardswish()
t = torch.tensor([-5.0, -2.0, 0.0, 2.0, 5.0])
res = layer(t)
print(res)
Результат выполнения кода:
tensor([0.0000, -0.3333, 0.0000, 2.0000, 5.0000])
Пример
Используем Hardswish в составе последовательной модели:
import torch
model = torch.nn.Sequential(
torch.nn.Linear(10, 20),
torch.nn.Hardswish(),
torch.nn.Linear(20, 5)
)
t = torch.randn(3, 10)
res = model(t)
print(res.shape)
Результат выполнения кода:
torch.Size([3, 5])
Пример
Применим Hardswish к двумерному тензору:
import torch
layer = torch.nn.Hardswish()
t = torch.tensor([
[-3.0, -1.0, 0.0],
[1.0, 3.0, 6.0]
])
res = layer(t)
print(res)
Результат выполнения кода:
tensor([
[0.0000, -0.3333, 0.0000],
[1.0000, 3.0000, 6.0000]
])
Смотрите также
-
класс
SiLU,
который реализует функцию активации swish -
класс
ReLU,
который реализует выпрямленную линейную функцию активации -
класс
Hardsigmoid,
который реализует кусочно-линейную аппроксимацию сигмоиды -
класс
Hardtanh,
который реализует кусочно-линейную функцию активации с ограничением