import torch import numpy as np import h5py import os import sys from pathlib import Path from tqdm import tqdm # Add project root to path project_root = str(Path(__file__).parent.parent) if project_root not in sys.path: sys.path.append(project_root) from src.models.student import LIPEV2Student def evaluate_on_file(model, h5_path, device, apply_coord_fix=True): results = { 'all': {'error': 0.0, 'count': 0}, 'frontal_45': {'error': 0.0, 'count': 0} } s_p, s_y = (-1, -1) if apply_coord_fix else (1, 1) if not os.path.exists(h5_path): return None with h5py.File(h5_path, 'r') as f: lp = f['left_patches'][:] rp = f['right_patches'][:] lm = f['landmarks'][:] g_gt = f['gaze'][:] with torch.no_grad(): for i in range(lp.shape[0]): # Inputs l_p = torch.from_numpy(lp[i]).float().unsqueeze(0).to(device) r_p = torch.from_numpy(rp[i]).float().unsqueeze(0).to(device) l_m = torch.from_numpy(lm[i]).float().view(1, -1).to(device) gt = torch.from_numpy(g_gt[i]).float().to(device) # Predict (Handles both DANN and non-DANN forward output count) out = model(l_p, l_m, state='A') p_l, y_l = out[0], out[1] out_r = model(r_p, l_m, state='A') p_r, y_r = out_r[0], out_r[1] def l2d(p, y): idx = torch.arange(90).float().to(device) pp, yp = torch.softmax(p, 1), torch.softmax(y, 1) return (torch.sum(pp*idx,1)*2-90), (torch.sum(yp*idx,1)*2-90) pl, yl = l2d(p_l, y_l) pr, yr = l2d(p_r, y_r) pf, yf = ((pl+pr)/2)*s_p, ((yl+yr)/2)*s_y gt_d = gt * (180.0/np.pi) error = (torch.abs(pf-gt_d[0]) + torch.abs(yf-gt_d[1])).item() results['all']['error'] += error results['all']['count'] += 1 if abs(gt_d[1].item()) <= 45.0: results['frontal_45']['error'] += error results['frontal_45']['count'] += 1 return {k: v['error']/(v['count']*2) for k, v in results.items() if v['count'] > 0} def run_comparison(model_path): device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') print(f"Comparing Preprocessing Methods for Model: {model_path}") # Load Model model = LIPEV2Student().to(device) model.load_state_dict(torch.load(model_path, map_location=device)) model.eval() datasets = { 'OLD (Raw Warp)': 'data/processed/gaze360_robust_v16_old.h5', 'NEW (Huy Filters)': 'data/processed/gaze360_robust_v16_new.h5' } print("\n" + "="*60) print(f"{'DATASET VERSION':<20} | {'ALL CASES MAE':<15} | {'FRONTAL 45 MAE':<15}") print("-"*60) for name, path in datasets.items(): res = evaluate_on_file(model, path, device) if res: print(f"{name:<20} | {res.get('all', 0):<15.4f} | {res.get('frontal_45', 0):<15.4f}") else: print(f"{name:<20} | {'FILE NOT FOUND':<33}") print("="*60) if __name__ == "__main__": # Test with the DANN Only pilot model best_model = 'checkpoints/dann_only/student_dann_final.pt' if os.path.exists(best_model): run_comparison(best_model) else: # Fallback to a LOPO checkpoint if DANN not ready print("DANN model not found, searching for LOPO checkpoint...") checkpoint_dir = 'checkpoints' pts = [f for f in os.listdir(checkpoint_dir) if f.endswith('.pt') and 'best' in f] if pts: run_comparison(os.path.join(checkpoint_dir, pts[0]))