import torch import torch.nn as nn import torch.optim as optim from torch.utils.data import DataLoader import os import time import numpy as np from tqdm import tqdm import sys import argparse from pathlib import Path # Add project root to path sys.path.append(str(Path(__file__).parent.parent)) from src.data.dataset import GazeDataset # --- 1. Loss Functions --- class SmoothAWLoss(nn.Module): def __init__(self, omega=8.0, alpha=1.5, theta=0.5, epsilon=1.0): super(SmoothAWLoss, self).__init__() self.omega = omega self.alpha = alpha self.theta = theta self.epsilon = epsilon def forward(self, y_pred, y_true): delta_y = (y_true - y_pred).abs() device = y_pred.device mask = delta_y < self.theta loss = torch.zeros_like(delta_y) loss[mask] = self.omega * torch.log(1 + torch.pow(delta_y[mask] / self.epsilon, self.alpha)) 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)) loss[~mask] = A * delta_y[~mask] - B return loss.mean() class DistillationLoss(nn.Module): def __init__(self, temperature=2.0, alpha=0.5, base_criterion=nn.MSELoss()): super(DistillationLoss, self).__init__() self.temp = temperature self.alpha = alpha self.base_criterion = base_criterion self.kl_div = nn.KLDivLoss(reduction='batchmean') def forward(self, student_logits, teacher_logits, labels): # student_logits/teacher_logits are tuples (pitch, yaw) p_s, y_s = student_logits p_t, y_t = teacher_logits # Soft targets loss loss_p_kd = self.kl_div(torch.log_softmax(p_s / self.temp, dim=1), torch.softmax(p_t / self.temp, dim=1)) * (self.temp ** 2) loss_y_kd = self.kl_div(torch.log_softmax(y_s / self.temp, dim=1), torch.softmax(y_t / self.temp, dim=1)) * (self.temp ** 2) # Hard targets loss (GT) loss_gt = self.base_criterion(student_logits, labels) # This assumes base_criterion handles the logits/labels correctly return self.alpha * (loss_p_kd + loss_y_kd) + (1 - self.alpha) * loss_gt # --- 2. Model Architecture with Ablation Toggles --- class FlexibleMiniConv(nn.Module): def __init__(self, pool_type='avg'): super(FlexibleMiniConv, self).__init__() self.conv = nn.Sequential( nn.Conv2d(1, 16, kernel_size=3, padding=0), nn.ReLU(inplace=True), nn.Conv2d(16, 32, kernel_size=3, padding=0), nn.ReLU(inplace=True), nn.Conv2d(32, 64, kernel_size=3, padding=0), nn.ReLU(inplace=True) ) self.pool_type = pool_type if pool_type == 'avg': self.pool = nn.AdaptiveAvgPool2d(1) self.out_channels = 64 else: # Dual Pool self.avg_pool = nn.AdaptiveAvgPool2d(1) self.max_pool = nn.AdaptiveMaxPool2d(1) self.out_channels = 128 def forward(self, x): x = self.conv(x) if self.pool_type == 'avg': return self.pool(x).flatten(1) else: return torch.cat([self.avg_pool(x), self.max_pool(x)], dim=1).flatten(1) class LIPEV2AblationModel(nn.Module): def __init__(self, pool_type='avg'): super(LIPEV2AblationModel, self).__init__() self.app_net = FlexibleMiniConv(pool_type=pool_type) # app_net output: out_channels * 4 patches app_out_dim = self.app_net.out_channels * 4 self.geo_net = nn.Sequential( nn.Linear(956, 256), nn.LayerNorm(256), nn.ReLU(inplace=True), nn.Linear(256, 256), nn.ReLU(inplace=True) ) self.post_concat_bn = nn.BatchNorm1d(app_out_dim + 256) self.fusion = nn.Sequential( nn.Linear(app_out_dim + 256, 256), nn.ReLU(inplace=True), nn.Dropout(0.1), nn.Linear(256, 128), nn.ReLU(inplace=True) ) self.pitch_head = nn.Linear(128, 90) self.yaw_head = nn.Linear(128, 90) def forward(self, patches, landmarks): batch_size = patches.shape[0] p_h, p_w = patches.shape[2], patches.shape[3] app_feat = self.app_net(patches.view(-1, 1, p_h, p_w)).view(batch_size, -1) geo_feat = self.geo_net(landmarks) combined = torch.cat([app_feat, geo_feat], dim=1) combined = self.post_concat_bn(combined) fused = self.fusion(combined) return self.pitch_head(fused), self.yaw_head(fused) # --- 3. Training Function --- def train_ablation(args): device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') print(f"\n>>> ABLATION STUDY ID: #{args.id} | Participant: {args.participant}") print(f">>> Config: Pool={args.pooling}, Loss={args.loss}, KD={args.kd}") # Load Data h5_dir = 'data/processed' train_files = [os.path.join(h5_dir, f) for f in os.listdir(h5_dir) if f.endswith('_v16.h5') and not f.startswith(args.participant)] val_files = [os.path.join(h5_dir, f) for f in os.listdir(h5_dir) if f.startswith(args.participant) and f.endswith('_v16.h5')] train_loader = DataLoader(GazeDataset(train_files, transform=True), batch_size=32, shuffle=True, num_workers=args.workers) val_loader = DataLoader(GazeDataset(val_files, transform=False), batch_size=32, shuffle=False, num_workers=args.workers) # Model model = LIPEV2AblationModel(pool_type=args.pooling).to(device) # KD Setup teacher_model = None if args.kd: from src.models.loss import GazeDistillationLoss # We assume the GazeDataset provides teacher predictions or we need a teacher model # For simplicity in this ablation, if GazeDataset has t_p, t_y, we use them kd_criterion = GazeDistillationLoss(w_aw=1.0, w_kd=0.5) print(">>> KD Enabled: Using GazeDistillationLoss") # Base Loss if args.loss == 'smoothaw': criterion = SmoothAWLoss() else: criterion = nn.MSELoss() optimizer = optim.AdamW(model.parameters(), lr=args.lr, weight_decay=1e-2) scheduler = optim.lr_scheduler.OneCycleLR(optimizer, max_lr=args.lr, steps_per_epoch=len(train_loader), epochs=args.epochs) idx_tensor = torch.arange(90).float().to(device) best_mae = float('inf') # Checkpoints dir ckpt_dir = f"checkpoints/ablation/id_{args.id}" os.makedirs(ckpt_dir, exist_ok=True) for epoch in range(1, args.epochs + 1): model.train() running_loss = 0.0 pbar = tqdm(train_loader, desc=f"Epoch {epoch}") for batch in pbar: # batch: [patches, landmarks, gaze, t_p, t_y, domain] batch = [b.to(device) for b in batch] optimizer.zero_grad() p_logits, y_logits = model(batch[0], batch[1]) s_p = (torch.softmax(p_logits, dim=1) * idx_tensor).sum(1) * 2 - 90 s_y = (torch.softmax(y_logits, dim=1) * idx_tensor).sum(1) * 2 - 90 if args.kd: # GazeDistillationLoss takes (p_logits, y_logits), gaze_gt, (t_p, t_y) loss, _, _ = kd_criterion((p_logits, y_logits), batch[2], (batch[3], batch[4])) else: # Simple loss for this ablation loss = criterion(torch.stack([s_p, s_y], dim=1), batch[2] * (180.0 / np.pi)) loss.backward() optimizer.step() scheduler.step() running_loss += loss.item() # Validation model.eval() total_error, count = 0.0, 0 with torch.no_grad(): for batch in val_loader: batch = [b.to(device) for b in batch] p_l, y_l = model(batch[0], batch[1]) p_d = (torch.softmax(p_l, dim=1) * idx_tensor).sum(1) * 2 - 90 y_d = (torch.softmax(y_l, dim=1) * idx_tensor).sum(1) * 2 - 90 gt_d = batch[2] * (180.0 / np.pi) total_error += (torch.abs(p_d - gt_d[:,0]) + torch.abs(y_d - gt_d[:,1])).sum().item() count += batch[2].shape[0] val_mae = total_error / (count * 2) if count > 0 else 0 print(f"Epoch {epoch}: Val MAE {val_mae:.4f}") if val_mae < best_mae: best_mae = val_mae torch.save(model.state_dict(), f"{ckpt_dir}/best_{args.participant}.pt") print(f"--- New Best: {best_mae:.4f} ---") # Final log with open('logs/ablation_results.log', 'a') as f: f.write(f"ID: {args.id}, Participant: {args.participant}, MAE: {best_mae:.4f}, Config: {args.pooling}/{args.loss}/KD={args.kd}\n") if __name__ == '__main__': parser = argparse.ArgumentParser() parser.add_argument('--id', type=int, required=True, help="Ablation ID (1-8)") parser.add_argument('--participant', type=str, default='p11') parser.add_argument('--pooling', choices=['avg', 'dual'], default='avg') parser.add_argument('--loss', choices=['mse', 'smoothaw'], default='mse') parser.add_argument('--kd', action='store_true') parser.add_argument('--epochs', type=int, default=50) parser.add_argument('--lr', type=float, default=1e-4) parser.add_argument('--workers', type=int, default=4) args = parser.parse_args() # Auto-config based on ID if provided (optional override) # 1: avg/mse/no, 2: dual/mse/no, 3: avg/saw/no, 4: avg/mse/yes # 5: dual/saw/no, 6: dual/mse/yes, 7: avg/saw/yes, 8: dual/saw/yes configs = { 1: ('avg', 'mse', False), 2: ('dual', 'mse', False), 3: ('avg', 'smoothaw', False), 4: ('avg', 'mse', True), 5: ('dual', 'smoothaw', False), 6: ('dual', 'mse', True), 7: ('avg', 'smoothaw', True), 8: ('dual', 'smoothaw', True), } if args.id in configs: args.pooling, args.loss, args.kd = configs[args.id] train_ablation(args)