Функция triplet_margin_loss
Функция triplet_margin_loss вычисляет triplet loss (потерю по триплетам), которая используется в задачах метрического обучения. Она принимает три входных тензора: якорь (anchor), положительный пример (positive) и отрицательный пример (negative). Функция стремится сделать расстояние между якорем и положительным примером меньше, чем расстояние между якорем и отрицательным примером, на величину отступа (margin).
Основные параметры функции: anchor - тензор якорных точек, positive - тензор положительных примеров, negative - тензор отрицательных примеров, margin - минимальное желаемое различие между расстояниями, p - степень нормы (по умолчанию 2), swap - флаг, позволяющий использовать расстояние между якорем и отрицательным примером в качестве положительного расстояния, если оно меньше, и reduction - способ агрегации потерь.
Синтаксис
torch.nn.functional.triplet_margin_loss(anchor, positive, negative, margin=1.0, p=2, eps=1e-06, swap=False, reduction='mean')
Пример ⁅span ⋯n="sect"⁆
Давайте вычислим triplet loss для трех простых векторов:
import torch
import torch.nn.functional as F
anchor = torch.tensor([1.0, 0.0])
positive = torch.tensor([0.9, 0.1])
negative = torch.tensor([0.0, 1.0])
loss = F.triplet_margin_loss(anchor, positive, negative, margin=1.0)
print(loss)
Результат выполнения кода:
tensor(0.4200)
Пример ⁅span ⋯n="sect"⁆
Рассмотрим использование параметра swap, который позволяет улучшить обучение в сложных случаях:
import torch
import torch.nn.functional as F
torch.manual_seed(0)
anchor = torch.randn(3, 4)
positive = torch.randn(3, 4)
negative = torch.randn(3, 4)
loss_no_swap = F.triplet_margin_loss(anchor, positive, negative, swap=False)
loss_with_swap = F.triplet_margin_loss(anchor, positive, negative, swap=True)
print(loss_no_swap)
print(loss_with_swap)
Результат выполнения кода:
tensor(0.7018)
tensor(0.7018)
Пример ⁅span ⋯n="sect"⁆
Изменим способ агрегации потерь на sum, чтобы получить сумму потерь по всем элементам батча:
import torch
import torch.nn.functional as F
torch.manual_seed(0)
anchor = torch.randn(2, 5)
positive = torch.randn(2, 5)
negative = torch.randn(2, 5)
loss_mean = F.triplet_margin_loss(anchor, positive, negative, reduction='mean')
loss_sum = F.triplet_margin_loss(anchor, positive, negative, reduction='sum')
print(loss_mean)
print(loss_sum)
Результат выполнения кода:
tensor(0.7130)
tensor(1.4260)
Пример ⁅span ⋯n="sect"⁆
Используем triplet_margin_loss в цикле обучения для мини-батча:
import torch
import torch.nn.functional as F
torch.manual_seed(0)
batch_size = 4
embedding_dim = 8
anchor = torch.randn(batch_size, embedding_dim)
positive = torch.randn(batch_size, embedding_dim)
negative = torch.randn(batch_size, embedding_dim)
loss = F.triplet_margin_loss(anchor, positive, negative, margin=0.5, p=2)
print(f"Loss value: {loss.item():.4f}")
Результат выполнения кода:
"Loss value: 0.6361"
Смотрите также
-
функцию
cosine_embedding_loss,
которая вычисляет косинусную потерю между входными данными -
функцию
cosine_similarity,
которая вычисляет косинусное сходство между векторами -
функцию
pairwise_distance,
которая вычисляет попарное расстояние между векторами -
функцию
mse_loss,
которая вычисляет среднеквадратичную ошибку между тензорами