File size: 6,192 Bytes
a10ba7f | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 | import torch
import torch.nn as nn
import torch.nn.functional as F
import numpy as np
class AdaptiveWingLoss(nn.Module):
def __init__(self, omega=14.0, theta=0.5, epsilon=1.0, alpha=2.1):
super(AdaptiveWingLoss, self).__init__()
self.omega = omega
self.theta = theta
self.epsilon = epsilon
self.alpha = alpha
def forward(self, y_pred, y_true):
"""
y_pred, y_true: (Batch, N)
"""
delta_y = (y_true - y_pred).abs()
device = y_pred.device
# Adaptive Wing Loss Formula
# L = omega * ln(1 + |delta_y / epsilon| ^ (alpha - y_true)) if delta_y < theta
# L = A * delta_y - B if delta_y >= theta
mask = delta_y < self.theta
# Part 1: delta_y < theta
loss1 = self.omega * torch.log(1 + torch.pow(delta_y[mask] / self.epsilon, self.alpha))
# Part 2: delta_y >= theta
theta_eps = torch.tensor(self.theta / self.epsilon, device=device)
A = self.omega * (1.0 / (1.0 + torch.pow(theta_eps, self.alpha))) * \
(self.alpha * torch.pow(theta_eps, self.alpha - 1.0) * (1.0 / self.epsilon))
B = A * self.theta - self.omega * torch.log(1.0 + torch.pow(theta_eps, self.alpha))
loss2 = A * delta_y[~mask] - B
return (loss1.sum() + loss2.sum()) / y_pred.numel()
class KDLoss(nn.Module):
def __init__(self, temperature=3.0):
super(KDLoss, self).__init__()
self.temperature = temperature
self.kl_div = nn.KLDivLoss(reduction='batchmean')
def forward(self, student_logits, teacher_logits):
"""
student_logits, teacher_logits: (Batch, NumBins)
"""
soft_log_probs = F.log_softmax(student_logits / self.temperature, dim=1)
soft_targets = F.softmax(teacher_logits / self.temperature, dim=1)
return self.kl_div(soft_log_probs, soft_targets) * (self.temperature ** 2)
class GazeDistillationLoss(nn.Module):
def __init__(self, w_aw=1.0, w_kd=0.5, temperature=3.0, focal_tau=4.2, focal_gamma=1.5):
super(GazeDistillationLoss, self).__init__()
self.aw_loss = AdaptiveWingLoss()
self.kd_loss = KDLoss(temperature=temperature)
self.w_aw = w_aw
self.w_kd = w_kd
self.focal_tau = focal_tau
self.focal_gamma = focal_gamma
def forward(self, student_outputs, gt_gaze, teacher_logits=None):
"""
student_outputs: Tuple(pitch_logits, yaw_logits) or Tuple(pitch, yaw, domain)
gt_gaze: (Batch, 2) in Radians or Degrees (depending on training)
teacher_logits: Tuple(t_pitch_logits, t_yaw_logits)
"""
# Handle 3 return values from DANN-enabled student
if len(student_outputs) == 3:
s_pitch_logits, s_yaw_logits, _ = student_outputs
else:
s_pitch_logits, s_yaw_logits = student_outputs
device = s_pitch_logits.device
# 1. Calculate continuous angles from student logits (expectation)
idx = torch.arange(90).float().to(device)
s_pitch_prob = F.softmax(s_pitch_logits, dim=1)
s_yaw_prob = F.softmax(s_yaw_logits, dim=1)
s_pitch_idx = torch.sum(s_pitch_prob * idx, dim=1)
s_yaw_idx = torch.sum(s_yaw_prob * idx, dim=1)
s_pitch_deg = s_pitch_idx * 2 - 90
s_yaw_deg = s_yaw_idx * 2 - 90
s_gaze = torch.stack([s_pitch_deg, s_yaw_deg], dim=1)
# Convert GT to degrees if it's in radians
gt_gaze_deg = gt_gaze * (180.0 / np.pi)
# --- Focal Weighting Logic ---
# Calculate angular error per sample: (abs_diff_p + abs_diff_y) / 2
with torch.no_grad():
errors = torch.abs(s_gaze - gt_gaze_deg).mean(dim=1)
# Weight = gamma if error > tau, else 1.0
weights = torch.where(errors >= self.focal_tau,
torch.tensor(self.focal_gamma, device=device),
torch.tensor(1.0, device=device))
# 2. AW Loss (Hard Label) - Manually apply weighting
delta_y = (gt_gaze_deg - s_gaze).abs()
# Part 1: delta_y < theta
mask = delta_y < self.aw_loss.theta
loss_val = torch.zeros_like(delta_y)
loss_val[mask] = self.aw_loss.omega * torch.log(1 + torch.pow(delta_y[mask] / self.aw_loss.epsilon, self.aw_loss.alpha))
# Part 2: delta_y >= theta
theta_eps = torch.tensor(self.aw_loss.theta / self.aw_loss.epsilon, device=device)
A = self.aw_loss.omega * (1.0 / (1.0 + torch.pow(theta_eps, self.aw_loss.alpha))) * \
(self.aw_loss.alpha * torch.pow(theta_eps, self.aw_loss.alpha - 1.0) * (1.0 / self.aw_loss.epsilon))
B = A * self.aw_loss.theta - self.aw_loss.omega * torch.log(1.0 + torch.pow(theta_eps, self.aw_loss.alpha))
loss_val[~mask] = A * delta_y[~mask] - B
# Apply Focal Weights to AW Loss
# weights is (Batch,), loss_val is (Batch, 2)
loss_aw = (loss_val.mean(dim=1) * weights).mean()
# 3. KD Loss (Soft Label)
loss_kd = 0
if teacher_logits is not None:
t_pitch_logits, t_yaw_logits = teacher_logits
# Apply weighting to KL Divergence
def weighted_kl(s_logits, t_logits, w):
soft_log_probs = F.log_softmax(s_logits / self.kd_loss.temperature, dim=1)
soft_targets = F.softmax(t_logits / self.kd_loss.temperature, dim=1)
# nn.KLDivLoss(reduction='none') returns (Batch, Bins)
kl = F.kl_div(soft_log_probs, soft_targets, reduction='none').sum(dim=1)
return (kl * w).mean() * (self.kd_loss.temperature ** 2)
loss_kd_pitch = weighted_kl(s_pitch_logits, t_pitch_logits, weights)
loss_kd_yaw = weighted_kl(s_yaw_logits, t_yaw_logits, weights)
loss_kd = (loss_kd_pitch + loss_kd_yaw) / 2
total_loss = self.w_aw * loss_aw + self.w_kd * loss_kd
return total_loss, loss_aw, loss_kd
|