| 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 |
|
|
| |
| sys.path.append(str(Path(__file__).parent.parent)) |
|
|
| from src.data.dataset import GazeDataset |
|
|
| |
|
|
| 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): |
| |
| p_s, y_s = student_logits |
| p_t, y_t = teacher_logits |
| |
| |
| 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) |
| |
| |
| loss_gt = self.base_criterion(student_logits, labels) |
| |
| return self.alpha * (loss_p_kd + loss_y_kd) + (1 - self.alpha) * loss_gt |
|
|
| |
|
|
| 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: |
| 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_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) |
|
|
| |
|
|
| 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}") |
|
|
| |
| 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 = LIPEV2AblationModel(pool_type=args.pooling).to(device) |
| |
| |
| teacher_model = None |
| if args.kd: |
| from src.models.loss import GazeDistillationLoss |
| |
| |
| kd_criterion = GazeDistillationLoss(w_aw=1.0, w_kd=0.5) |
| print(">>> KD Enabled: Using GazeDistillationLoss") |
|
|
| |
| 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') |
| |
| |
| 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 = [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: |
| |
| loss, _, _ = kd_criterion((p_logits, y_logits), batch[2], (batch[3], batch[4])) |
| else: |
| |
| 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() |
|
|
| |
| 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} ---") |
|
|
| |
| 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() |
| |
| |
| |
| |
| 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) |
|
|