Download src/data/check_coordinate_alignment.py from thanhhuyvan/Gaze-LIPE: direct link, hf CLI and curl.
- Browser
- Download file 2.65 kB
-
https://huggingface.co/thanhhuyvan/Gaze-LIPE/resolve/main/src/data/check_coordinate_alignment.py
- Command line
-
hf download hf://thanhhuyvan/Gaze-LIPE/src/data/check_coordinate_alignment.py
-
curl -L -o check_coordinate_alignment.py https://huggingface.co/thanhhuyvan/Gaze-LIPE/resolve/main/src/data/check_coordinate_alignment.py
2.65 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.parent) | |
| if project_root not in sys.path: | |
| sys.path.append(project_root) | |
| from src.models.student import LIPEV2Student | |
| def check_alignment(model_path, h5_path): | |
| device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') | |
| model = LIPEV2Student().to(device) | |
| state_dict = torch.load(model_path, map_location=device) | |
| model.load_state_dict(state_dict) | |
| model.eval() | |
| # Metrics for different combinations | |
| # (P_sign, Y_sign) | |
| combinations = [ | |
| (1, 1), # Normal | |
| (1, -1), # Inverted Yaw | |
| (-1, 1), # Inverted Pitch | |
| (-1, -1) # Both Inverted | |
| ] | |
| errors = {c: 0.0 for c in combinations} | |
| count = 0 | |
| with h5py.File(h5_path, 'r') as f: | |
| lp = f['left_patches'][:] | |
| rp = f['right_patches'][:] | |
| lm = f['landmarks'][:] | |
| g_gt = f['gaze'][:] | |
| num_samples = min(500, lp.shape[0]) # Use subset for speed | |
| with torch.no_grad(): | |
| for i in range(num_samples): | |
| p_l, y_l = model(torch.from_numpy(lp[i]).float().unsqueeze(0).to(device), | |
| torch.from_numpy(lm[i]).float().view(1, -1).to(device), state='A') | |
| p_r, y_r = model(torch.from_numpy(rp[i]).float().unsqueeze(0).to(device), | |
| torch.from_numpy(lm[i]).float().view(1, -1).to(device), state='A') | |
| 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, (yl+yr)/2 | |
| gt_d = torch.from_numpy(g_gt[i]).to(device) * (180.0/np.pi) | |
| for ps, ys in combinations: | |
| err = (torch.abs(ps*pf - gt_d[0]) + torch.abs(ys*yf - gt_d[1])).item() | |
| errors[(ps, ys)] += err | |
| count += 1 | |
| print(f"\n" + "="*40) | |
| print(f"{'SIGN (Pitch, Yaw)':<20} | {'MAE (deg)':<10}") | |
| print(f"-"*40) | |
| for c, err in errors.items(): | |
| mae = err / (count * 2) | |
| print(f"{str(c):<20} | {mae:.4f}") | |
| print(f"="*40) | |
| if __name__ == "__main__": | |
| check_alignment('checkpoints/best_student_p04.pt', 'data/processed/gaze360_robust_v16.h5') | |