Функция F.elu
Функция F.elu применяет экспоненциальную линейную активацию
(Exponential Linear Unit) к каждому элементу входного тензора.
Эта функция была предложена как улучшение ReLU: для отрицательных
значений она не обнуляет их, а делает их близкими к нулю с помощью
экспоненциальной функции, что сохраняет информацию и ускоряет
сходимость.
В качестве первого параметра функция принимает тензор input.
Вторым параметром передаётся коэффициент alpha (по умолчанию 1.0),
который управляет насыщением для отрицательных значений.
Третьим параметром можно указать inplace (по умолчанию False),
чтобы выполнить операцию на месте.
Синтаксис
torch.nn.functional.elu(input, alpha=1.0, inplace=False)
Пример
Давайте применим функцию ELU к простому тензору с положительными и отрицательными числами:
import torch
import torch.nn.functional as F
t = torch.tensor([-3.0, -1.0, 0.0, 1.0, 3.0])
res = F.elu(t)
print(res)
Результат выполнения кода:
tensor([-0.9502, -0.6321, 0.0000, 1.0000, 3.0000])
Пример
Теперь изменим параметр alpha, чтобы сделать отрицательную
часть более крутой:
import torch
import torch.nn.functional as F
t = torch.tensor([-3.0, -1.0, 0.0, 1.0, 3.0])
res = F.elu(t, alpha=2.0)
print(res)
Результат выполнения кода:
tensor([-1.9004, -1.2642, 0.0000, 1.0000, 3.0000])
Пример
Используем функцию в составе нейросети, добавив её после линейного слоя для активации выходов:
import torch
import torch.nn as nn
import torch.nn.functional as F
class MyNet(nn.Module):
def __init__(self):
super().__init__()
self.fc = nn.Linear(5, 3)
def forward(self, x):
x = self.fc(x)
x = F.elu(x, alpha=0.5)
return x
torch.manual_seed(0)
model = MyNet()
inp = torch.randn(2, 5)
out = model(inp)
print(out)
Результат выполнения кода:
tensor([
[-0.0179, 0.1075, 0.2175],
[-0.3445, -0.1158, 0.2488]
])
Смотрите также
-
функцию
relu,
которая обнуляет отрицательные значения, сохраняя положительные -
функцию
leaky_relu,
которая для отрицательных значений использует линейный коэффициент -
функцию
selu,
которая является масштабированной версией ELU с самоподобием -
функцию
celu,
которая обобщает ELU для использования в классификации