GPUCodeForces/S1/wut0n_#7/tripletloss_torchcode.py

49 lines
1.6 KiB
Python

import torch
import torch.nn as nn
class Model(nn.Module):
"""
Triplet Loss implementation - A loss function for deep metric learning.
Enforces that anchor samples are closer to positive samples than to negative samples.
"""
def __init__(self, margin=1.0, p=2.0, eps=1e-6):
super(Model, self).__init__()
self.margin = margin
self.p = p
self.eps = eps
def forward(self, anchor: torch.Tensor, positive: torch.Tensor, negative: torch.Tensor) -> torch.Tensor:
"""
Compute triplet loss.
Args:
anchor (torch.Tensor): Anchor samples [batch_size, feature_dim]
positive (torch.Tensor): Positive samples [batch_size, feature_dim]
negative (torch.Tensor): Negative samples [batch_size, feature_dim]
Returns:
torch.Tensor: Scalar triplet loss value
"""
# Compute distances
dist_positive = torch.pairwise_distance(anchor, positive, p=self.p, eps=self.eps)
dist_negative = torch.pairwise_distance(anchor, negative, p=self.p, eps=self.eps)
# Compute triplet loss
losses = torch.relu(dist_positive - dist_negative + self.margin)
return torch.mean(losses)
batch_size = 256
feature_dim = 512
margin = 1.0
def get_inputs():
# Generate triplet data
anchor = torch.randn(batch_size, feature_dim)
positive = torch.randn(batch_size, feature_dim)
negative = torch.randn(batch_size, feature_dim)
return [anchor, positive, negative]
def get_init_inputs():
return [margin]