Функция F.softmax
Функция F.softmax применяется к тензору и возвращает новый тензор с теми же размерами, где каждое значение преобразовано в вероятность. Первым параметром функция принимает входной тензор, вторым параметром - индекс оси, вдоль которой выполняется нормализация. Также доступен параметр dtype для задания типа данных выходного тензора.
Синтаксис
torch.nn.functional.softmax(input, dim, [dtype])
Пример
Давайте преобразуем одномерный тензор в распределение вероятностей:
import torch
import torch.nn.functional as F
t = torch.tensor([1.0, 2.0, 3.0])
res = F.softmax(t, dim=0)
print(res)
Результат выполнения кода:
tensor([0.0900, 0.2447, 0.6652])
Пример
Выполним softmax для двумерного тензора по разным осям. Сначала по строкам (dim=1):
import torch
import torch.nn.functional as F
t = torch.tensor([[1.0, 2.0], [3.0, 4.0]])
res = F.softmax(t, dim=1)
print(res)
Результат выполнения кода:
tensor([
[0.2689, 0.7311],
[0.2689, 0.7311],
])
А теперь по столбцам (dim=0):
import torch
import torch.nn.functional as F
t = torch.tensor([[1.0, 2.0], [3.0, 4.0]])
res = F.softmax(t, dim=0)
print(res)
Результат выполнения кода:
tensor([
[0.1192, 0.1192],
[0.8808, 0.8808],
])
Пример
Применим softmax с указанием типа данных для повышения стабильности вычислений:
import torch
import torch.nn.functional as F
t = torch.tensor([1.0, 2.0, 3.0])
res = F.softmax(t, dim=0, dtype=torch.float64)
print(res)
print(res.dtype)
Результат выполнения кода:
tensor([0.0900, 0.2447, 0.6652], dtype=torch.float64)
torch.float64
Пример
Используем softmax для получения вероятностей классов в задаче классификации:
import torch
import torch.nn.functional as F
logits = torch.tensor([0.2, 1.0, 0.5])
probs = F.softmax(logits, dim=0)
print(probs)
print(probs.sum())
Результат выполнения кода:
tensor([0.2016, 0.4496, 0.3488])
tensor(1.0000)
Смотрите также
-
функцию
log_softmax,
которая вычисляет логарифм от softmax -
функцию
softmin,
которая аналогична softmax, но с обратным знаком -
функцию
gumbel_softmax,
которая позволяет делать дискретный выбор с дифференцируемостью -
функцию
cross_entropy,
которая сочетает в себе softmax и логарифмическую потерю