Download src/utils/analyze_errors.py from thanhhuyvan/Gaze-LIPE: direct link, hf CLI and curl.
- Browser
- Download file 2.88 kB
-
https://huggingface.co/thanhhuyvan/Gaze-LIPE/resolve/main/src/utils/analyze_errors.py
- Command line
-
hf download hf://thanhhuyvan/Gaze-LIPE/src/utils/analyze_errors.py
-
curl -L -o analyze_errors.py https://huggingface.co/thanhhuyvan/Gaze-LIPE/resolve/main/src/utils/analyze_errors.py
2.88 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 | |
| sys.path.append(str(Path(__file__).parent.parent.parent)) | |
| from src.models.student import LIPEV2Student | |
| def analyze_worst_errors(model_path, h5_path, top_n=10): | |
| device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') | |
| print(f"Analyzing errors for model: {model_path}") | |
| model = LIPEV2Student().to(device) | |
| state_dict = torch.load(model_path, map_location=device) | |
| model.load_state_dict(state_dict, strict=False) | |
| model.eval() | |
| errors = [] | |
| with h5py.File(h5_path, 'r') as f: | |
| lp = torch.from_numpy(f['left_patches'][:]).float().to(device) | |
| rp = torch.from_numpy(f['right_patches'][:]).float().to(device) | |
| lm = torch.from_numpy(f['landmarks'][:]).float().view(lp.shape[0], -1).to(device) | |
| # Average left and right gaze for ground truth | |
| gt_l = f['left_gaze'][:] | |
| gt_r = f['right_gaze'][:] | |
| gt = torch.from_numpy((gt_l + gt_r) / 2).float().to(device) | |
| with torch.no_grad(): | |
| p_logits_l, y_logits_l, _ = model(lp, lm, state='A') | |
| p_logits_r, y_logits_r, _ = model(rp, lm, state='A') | |
| # Logits to deg | |
| def logits_to_deg(p, y): | |
| idx = torch.arange(90).float().to(device) | |
| p_deg = (torch.softmax(p, dim=1) @ idx) * 2 - 90 | |
| y_deg = (torch.softmax(y, dim=1) @ idx) * 2 - 90 | |
| return p_deg, y_deg | |
| p_l, y_l = logits_to_deg(p_logits_l, y_logits_l) | |
| p_r, y_r = logits_to_deg(p_logits_r, y_logits_r) | |
| p_pred = (p_l + p_r) / 2 | |
| y_pred = (y_l + y_r) / 2 | |
| gt_deg = gt * (180.0 / np.pi) | |
| # Error per sample | |
| sample_errors = (torch.abs(p_pred - gt_deg[:, 0]) + torch.abs(y_pred - gt_deg[:, 1])).cpu().numpy() | |
| for i in range(len(sample_errors)): | |
| errors.append({ | |
| 'idx': i, | |
| 'mae': sample_errors[i] / 2, | |
| 'gt_pitch': gt_deg[i, 0].item(), | |
| 'gt_yaw': gt_deg[i, 1].item(), | |
| 'pred_pitch': p_pred[i].item(), | |
| 'pred_yaw': y_pred[i].item() | |
| }) | |
| # Sort by error descending | |
| errors.sort(key=lambda x: x['mae'], reverse=True) | |
| print(f"\nTop {top_n} Worst Errors:") | |
| print(f"{'Idx':<8} | {'MAE':<10} | {'GT (P,Y)':<20} | {'Pred (P,Y)':<20}") | |
| print("-" * 70) | |
| for e in errors[:top_n]: | |
| print(f"{e['idx']:<8} | {e['mae']:.4f} | ({e['gt_pitch']:.1f}, {e['gt_yaw']:.1f}) | ({e['pred_pitch']:.1f}, {e['pred_yaw']:.1f})") | |
| if __name__ == "__main__": | |
| analyze_worst_errors( | |
| model_path='checkpoints/best_student_p08.pt', | |
| h5_path='data/processed/p08_v16_new.h5' | |
| ) | |