Класс PReLU
Класс PReLU (Parametric Rectified Linear Unit) применяется к данным как функция активации.
В отличие от ReLU, которая обнуляет отрицательные значения, PReLU умножает их на обучаемый параметр a.
Это позволяет сети адаптивно регулировать наклон для отрицательных входных данных.
Первым параметром конструктор принимает количество каналов num_parameters (по умолчанию 1),
вторым - начальное значение init (по умолчанию 0.25).
Синтаксис
torch.nn.PReLU(num_parameters=1, init=0.25)
Параметры
-
num_parameters(int) - количество обучаемых параметровa. Обычно равно числу каналов входного тензора или1для общих параметров. По умолчанию1. -
init(float) - начальное значение параметраa. По умолчанию0.25.
Пример
Давайте создадим слой PReLU с одним общим параметром и применим его к тензору:
import torch
import torch.nn as nn
prelu = nn.PReLU()
t = torch.tensor([-2.0, -1.0, 0.0, 1.0, 2.0])
res = prelu(t)
print(res)
Результат выполнения кода:
tensor([-0.5000, -0.2500, 0.0000, 1.0000, 2.0000], grad_fn=<PReLUBackward0>)
Отрицательные значения умножаются на начальный параметр 0.25.
Пример
Теперь создадим PReLU с параметрами для каждого канала и посмотрим на обучаемый параметр:
import torch
import torch.nn as nn
prelu = nn.PReLU(num_parameters=3)
for name, param in prelu.named_parameters():
print(f"{name}: {param.data}")
Результат выполнения кода:
weight: tensor([0.2500, 0.2500, 0.2500])
Все три параметра инициализированы значением 0.25.
Пример
Применим слой к многоканальному тензору и обновим параметр в процессе обучения:
import torch
import torch.nn as nn
torch.manual_seed(0)
prelu = nn.PReLU(num_parameters=2)
t = torch.randn(1, 2, 3)
res = prelu(t)
print("Output:", res)
print("Parameter:", prelu.weight.data)
# Simulate a backward pass
loss = res.sum()
loss.backward()
with torch.no_grad():
prelu.weight -= 0.1 * prelu.weight.grad
print("Updated parameter:", prelu.weight.data)
Результат выполнения кода:
Output: tensor([[[-0.0879, 0.8117, -0.1250],
[ 0.3895, -0.5482, 0.3447]]], grad_fn=<PReLUBackward0>)
Parameter: tensor([0.2500, 0.2500])
Updated parameter: tensor([0.2104, 0.2224])
Параметр обновился после градиентного шага.