forked from ccf-ai-infra/GPUCodeForces
49 lines
1.6 KiB
Python
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]
|