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/1.13/generated/torch.nn.TripletMarginWithDistanceLoss.html