Gaze-LIPE / scripts /compare_preprocessing.py
thanhhuyvan's picture
Initial release of LIPE V2 GOLD
a10ba7f
Raw
History Blame Contribute Delete
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]))