TripletMarginWithDistanceLoss
-
class torch.nn.TripletMarginWithDistanceLoss(*, distance_function=None, margin=1.0, swap=False, reduction='mean')[source] -
Создаёт критерий, который измеряет потерю триплета, задавая входные тензоры , и (представляющие якорь, положительный и отрицательный примеры соответственно) и неотрицательную вещественную функцию («функцию расстояния»), используемую для вычисления отношения между якорем и положительным примером («положительное расстояние») и якорем и отрицательным примером («отрицательное расстояние»).
Несведённая потеря (т.е., с
reductionустановленным в'none') может быть описана как:где — размер пакета; — неотрицательная вещественная функция, количественно определяющая близость двух тензоров, называемая
distance_function; и — неотрицательный марж, представляющий минимальную разницу между положительными и отрицательными расстояниями, необходимую для того, чтобы потеря была равна 0. Входные тензоры имеют элементов каждый и могут иметь любую форму, которую может обработать функция расстояния.Если
reductionне'none'(по умолчанию'mean'), то:См. также
TripletMarginLoss, которое вычисляет потерю триплета для входных тензоров, используя расстояние в качестве функции расстояния.- Параметры
-
-
distance_function (Callable, optional) – Неотрицательная вещественная функция, количественно определяющая близость двух тензоров. Если не указано, будет использоваться
nn.PairwiseDistance. Значение по умолчанию:None - margin (float, optional) – Неотрицательный марж, представляющий собой минимальную разницу между положительными и отрицательными расстояниями, необходимую для того, чтобы потеря была равна 0. Более крупные маржи наказывают случаи, когда отрицательные примеры не достаточно удалены от якорей по отношению к положительным примерам. Значение по умолчанию: .
-
swap (bool, optional) – Необходимо ли использовать обмен расстояниями, описанный в статье
Learning shallow convolutional feature descriptors with triplet lossesВ. Балнтасом, Е. Рибой и др. Если True, и если положительный пример ближе к отрицательному примеру, чем якорь, меняет положительный пример и якорь в вычислении потери. Значение по умолчанию:False. -
reduction (str, optional) – Указывает (необязательное) сведение, применяемое к выходу:
'none'|'mean'|'sum'.'none': не будет применено сведение,'mean': сумма результата будет разделена на число элементов в результате,'sum': результат будет суммирован. Значение по умолчанию:'mean'
-
distance_function (Callable, optional) – Неотрицательная вещественная функция, количественно определяющая близость двух тензоров. Если не указано, будет использоваться
- Форма:
-
- Вход: , где представляет любое число дополнительных измерений, поддерживаемых функцией расстояния.
- Выход: тензор формы , если
reductionравно'none', или скаляр в противном случае.
Примеры:
>>> # Initialize embeddings >>> embedding = nn.Embedding(1000, 128) >>> anchor_ids = torch.randint(0, 1000, (1,)) >>> positive_ids = torch.randint(0, 1000, (1,)) >>> negative_ids = torch.randint(0, 1000, (1,)) >>> anchor = embedding(anchor_ids) >>> positive = embedding(positive_ids) >>> negative = embedding(negative_ids) >>> >>> # Built-in Distance Function >>> triplet_loss = \ >>> nn.TripletMarginWithDistanceLoss(distance_function=nn.PairwiseDistance()) >>> output = triplet_loss(anchor, positive, negative) >>> output.backward() >>> >>> # Custom Distance Function >>> def l_infinity(x1, x2): >>> return torch.max(torch.abs(x1 - x2), dim=1).values >>> >>> triplet_loss = ( >>> nn.TripletMarginWithDistanceLoss(distance_function=l_infinity, margin=1.5)) >>> output = triplet_loss(anchor, positive, negative) >>> output.backward() >>> >>> # Custom Distance Function (Lambda) >>> triplet_loss = ( >>> nn.TripletMarginWithDistanceLoss( >>> distance_function=lambda x, y: 1.0 - F.cosine_similarity(x, y))) >>> output = triplet_loss(anchor, positive, negative) >>> output.backward()
- Ссылка:
-
В. Балнтас и др.: Обучение неглубоких сверточных описателей признаков с потерями триплета: http://www.bmva.org/bmvc/2016/papers/paper119/index.html
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/2.1/generated/torch.nn.TripletMarginWithDistanceLoss.html