Класс TripletMarginLoss
Класс TripletMarginLoss реализует triplet loss,
которая используется при обучении моделей для задач
метрического обучения и поиска. Функция принимает
три тензора: якорь (anchor), положительный пример
(positive) и отрицательный пример (negative).
Основная идея заключается в том, чтобы расстояние
между якорем и положительным примером было меньше,
чем расстояние между якорем и отрицательным примером,
на заданную величину отступа (margin).
Основные параметры класса:
-
margin- отступ, минимальное значение, на которое расстояние до отрицательного примера должно превышать расстояние до положительного -
p- норма для вычисления расстояния (1 - манхэттенское, 2 - евклидово) -
swap- еслиTrue, то для отрицательного примера используется минимальное расстояние до якоря или положительного примера -
reduction- способ агрегации потерь (none,mean,sum)
Синтаксис
torch.nn.TripletMarginLoss(
margin=1.0,
p=2.0,
eps=1e-06,
swap=False,
reduction='mean'
)
Пример
Давайте рассмотрим базовый пример использования класса
TripletMarginLoss для трех простых тензоров:
import torch
import torch.nn as nn
criterion = nn.TripletMarginLoss(margin=1.0)
anchor = torch.tensor([
[1.0, 0.0],
[0.0, 1.0]
])
positive = torch.tensor([
[0.8, 0.1],
[0.1, 0.9]
])
negative = torch.tensor([
[-1.0, 0.0],
[0.0, -1.0]
])
loss = criterion(anchor, positive, negative)
print(loss)
Результат выполнения кода:
tensor(0.0000)
Пример
Теперь рассмотрим случай, когда положительный пример находится далеко от якоря, а отрицательный - близко. В этом случае функция потерь вернет положительное значение:
import torch
import torch.nn as nn
criterion = nn.TripletMarginLoss(margin=1.0)
anchor = torch.tensor([
[0.0, 0.0],
[1.0, 1.0]
])
positive = torch.tensor([
[0.9, 0.9],
[0.0, 0.0]
])
negative = torch.tensor([
[0.1, 0.1],
[0.9, 0.9]
])
loss = criterion(anchor, positive, negative)
print(loss)
Результат выполнения кода:
tensor(1.2800)
Пример
Используем параметр swap для более гибкого
вычисления потерь. При swap=True расстояние до
отрицательного примера вычисляется как минимальное
расстояние до якоря или положительного примера:
import torch
import torch.nn as nn
criterion = nn.TripletMarginLoss(
margin=1.0,
swap=True
)
torch.manual_seed(0)
anchor = torch.randn(3, 5)
positive = torch.randn(3, 5)
negative = torch.randn(3, 5)
loss = criterion(anchor, positive, negative)
print(loss)
Результат выполнения кода:
tensor(1.2589)
Пример
Используем разные значения нормы p для вычисления
расстояния. При p=1 используется манхэттенское
расстояние, при p=2 - евклидово:
import torch
import torch.nn as nn
criterion_l1 = nn.TripletMarginLoss(
margin=1.0,
p=1.0
)
criterion_l2 = nn.TripletMarginLoss(
margin=1.0,
p=2.0
)
torch.manual_seed(0)
anchor = torch.randn(2, 4)
positive = torch.randn(2, 4)
negative = torch.randn(2, 4)
loss_l1 = criterion_l1(anchor, positive, negative)
loss_l2 = criterion_l2(anchor, positive, negative)
print(loss_l1, loss_l2)
Результат выполнения кода:
tensor(0.6403) tensor(0.3911)
Пример
Рассмотрим использование параметра reduction,
который управляет способом агрегации потерь.
При reduction='none' возвращается тензор
индивидуальных потерь для каждого элемента:
import torch
import torch.nn as nn
criterion_none = nn.TripletMarginLoss(
margin=1.0,
reduction='none'
)
criterion_mean = nn.TripletMarginLoss(
margin=1.0,
reduction='mean'
)
criterion_sum = nn.TripletMarginLoss(
margin=1.0,
reduction='sum'
)
torch.manual_seed(0)
anchor = torch.randn(3, 4)
positive = torch.randn(3, 4)
negative = torch.randn(3, 4)
loss_none = criterion_none(anchor, positive, negative)
loss_mean = criterion_mean(anchor, positive, negative)
loss_sum = criterion_sum(anchor, positive, negative)
print(loss_none)
print(loss_mean)
print(loss_sum)
Результат выполнения кода:
tensor([1.3244, 0.0000, 1.4024])
tensor(0.9089)
tensor(2.7268)
Пример
Используем класс TripletMarginLoss в процессе
обучения простой модели для задачи поиска похожих
объектов:
import torch
import torch.nn as nn
import torch.optim as optim
class SimpleModel(nn.Module):
def __init__(self):
super().__init__()
self.fc = nn.Linear(10, 5)
def forward(self, x):
return self.fc(x)
model = SimpleModel()
criterion = nn.TripletMarginLoss(margin=1.0)
optimizer = optim.SGD(model.parameters(), lr=0.01)
torch.manual_seed(0)
anchor = torch.randn(4, 10)
positive = torch.randn(4, 10)
negative = torch.randn(4, 10)
anchor_emb = model(anchor)
positive_emb = model(positive)
negative_emb = model(negative)
loss = criterion(anchor_emb, positive_emb, negative_emb)
optimizer.zero_grad()
loss.backward()
optimizer.step()
print(loss.item())
Результат выполнения кода:
1.0698050260543823
Смотрите также
-
класс
CrossEntropyLoss,
который используется для задач классификации -
класс
CosineEmbeddingLoss,
который использует косинусное расстояние между эмбеддингами -
класс
MSELoss,
который вычисляет среднеквадратичную ошибку -
класс
MarginRankingLoss,
который используется для ранжирования с отступом