РЕПЕТИТОР математика физика информатика
Для школьников и студентов. Подтягивание пробелов. ЦЭ, ЦТ, ОГЭ, ЕГЭ.
Идет набор на ЛЕТО. Жмите для подробностей:)
442 of 769 menu

Функция 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,
    которая вычисляет среднеквадратичную ошибку между тензорами
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить