Функция F.dropout
Функция F.dropout применяет регуляризацию dropout к входному тензору.
Она случайным образом обнуляет элементы тензора с заданной вероятностью p
и масштабирует оставшиеся элементы для сохранения суммы активаций.
Первый параметр функции - входной тензор, второй - вероятность отбрасывания.
Третий параметр training указывает, применяется ли dropout во время обучения.
Четвертый параметр inplace позволяет изменять тензор на месте.
Синтаксис
torch.nn.functional.dropout(input, p, training, inplace)
Пример
Давайте применим dropout с вероятностью 0.5 к простому тензору в режиме обучения:
import torch
import torch.nn.functional as F
torch.manual_seed(0)
t = torch.tensor([1.0, 2.0, 3.0, 4.0, 5.0])
res = F.dropout(t, p=0.5, training=True)
print(res)
Результат выполнения кода:
tensor([0., 0., 6., 0., 10.])
Пример
В режиме оценки F.dropout не изменяет тензор, что используется при валидации и тестировании модели:
import torch
import torch.nn.functional as F
torch.manual_seed(0)
t = torch.tensor([1.0, 2.0, 3.0, 4.0, 5.0])
res = F.dropout(t, p=0.5, training=False)
print(res)
Результат выполнения кода:
tensor([1., 2., 3., 4., 5.])
Пример
Параметр inplace позволяет изменять тензор на месте, что экономит память:
import torch
import torch.nn.functional as F
torch.manual_seed(42)
t = torch.tensor([1.0, 2.0, 3.0, 4.0])
res = F.dropout(t, p=0.3, training=True, inplace=True)
print(res)
Результат выполнения кода:
tensor([0., 0., 4.2857, 5.7143])
Пример
Применим dropout к двумерному тензору, обрабатывая каждый элемент независимо:
import torch
import torch.nn.functional as F
torch.manual_seed(1)
t = torch.tensor([
[1.0, 2.0, 3.0],
[4.0, 5.0, 6.0],
])
res = F.dropout(t, p=0.3, training=True)
print(res)
Результат выполнения кода:
tensor([
[ 0., 0., 4.2857],
[5.7143, 0., 8.5714],
])
Смотрите также
-
функцию
dropout2d,
которая применяет dropout к каналам изображений -
функцию
alpha_dropout,
которая сохраняет среднее значение и дисперсию входных данных -
функцию
relu,
которая применяет функцию активации ReLU к тензору -
модуль
batch_norm,
который нормализует данные по батчам