Функция F.one_hot
Функция F.one_hot преобразует тензор, содержащий целочисленные метки классов, в тензор one-hot кодирования. Первым параметром функция принимает входной тензор с индексами классов. Вторым параметром можно указать общее количество классов через num_classes. Результатом является тензор, где каждая строка содержит единицу на позиции, соответствующей индексу класса, и нули на остальных позициях.
Синтаксис
torch.nn.functional.one_hot(tensor, num_classes=-1)
Пример
Давайте преобразуем тензор с метками классов в one-hot представление:
import torch
import torch.nn.functional as F
t = torch.tensor([0, 1, 2, 3])
res = F.one_hot(t)
print(res)
Результат выполнения кода:
tensor([
[1, 0, 0, 0],
[0, 1, 0, 0],
[0, 0, 1, 0],
[0, 0, 0, 1],
])
Пример
Укажем количество классов вручную с помощью параметра num_classes:
import torch
import torch.nn.functional as F
t = torch.tensor([1, 2, 0])
res = F.one_hot(t, num_classes=5)
print(res)
Результат выполнения кода:
tensor([
[0, 1, 0, 0, 0],
[0, 0, 1, 0, 0],
[1, 0, 0, 0, 0],
])
Пример
Применяем one-hot кодирование к двумерному тензору пакета данных:
import torch
import torch.nn.functional as F
t = torch.tensor([
[1, 0, 3],
[2, 1, 0],
])
res = F.one_hot(t, num_classes=4)
print(res.shape)
print(res)
Результат выполнения кода:
torch.Size([2, 3, 4])
tensor([
[
[0, 1, 0, 0],
[1, 0, 0, 0],
[0, 0, 0, 1],
],
[
[0, 0, 1, 0],
[0, 1, 0, 0],
[1, 0, 0, 0],
],
])
Смотрите также
-
функцию
embedding,
которая выполняет поиск по таблице вложений -
функцию
softmax,
которая преобразует логиты в вероятности классов -
функцию
cross_entropy,
которая использует метки классов для вычисления ошибки -
функцию
nll_loss,
которая вычисляет отрицательное логарифмическое правдоподобие