Download scripts/compare_preprocessing.py from thanhhuyvan/Gaze-LIPE: direct link, hf CLI and curl.
- Browser
- Download file 3.82 kB
-
https://huggingface.co/thanhhuyvan/Gaze-LIPE/resolve/main/scripts/compare_preprocessing.py
- Command line
-
hf download hf://thanhhuyvan/Gaze-LIPE/scripts/compare_preprocessing.py
-
curl -L -o compare_preprocessing.py https://huggingface.co/thanhhuyvan/Gaze-LIPE/resolve/main/scripts/compare_preprocessing.py
3.82 kB
| 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])) | |