| import h5py |
| import numpy as np |
| import torch |
| import torch.nn.functional as F |
| import sys |
| import os |
|
|
| def check_teacher_accuracy(h5_path): |
| with h5py.File(h5_path, 'r') as f: |
| if 'teacher_pitch_logits' not in f: |
| print(f"No teacher labels in {h5_path}") |
| return |
| |
| p_logits = torch.from_numpy(f['teacher_pitch_logits'][:]) |
| y_logits = torch.from_numpy(f['teacher_yaw_logits'][:]) |
| gt_l_gaze = f['left_gaze'][:] |
| gt_r_gaze = f['right_gaze'][:] |
| |
| |
| gt_gaze = (gt_l_gaze + gt_r_gaze) / 2 |
| gt_gaze_deg = gt_gaze * (180.0 / np.pi) |
| |
| |
| idx = torch.arange(90).float() |
| p_prob = F.softmax(p_logits, dim=1) |
| y_prob = F.softmax(y_logits, dim=1) |
| |
| |
| p_deg1 = torch.sum(p_prob * idx, dim=1) * 4 - 180 |
| y_deg1 = torch.sum(y_prob * idx, dim=1) * 4 - 180 |
| mae1 = (torch.abs(p_deg1 - torch.from_numpy(gt_gaze_deg[:, 0])).mean() + |
| torch.abs(y_deg1 - torch.from_numpy(gt_gaze_deg[:, 1])).mean()) / 2 |
| print(f"MAE with idx*4 - 180: {mae1:.4f}") |
|
|
| |
| p_deg2 = torch.sum(p_prob * idx, dim=1) * 2 - 90 |
| y_deg2 = torch.sum(y_prob * idx, dim=1) * 2 - 90 |
| mae2 = (torch.abs(p_deg2 - torch.from_numpy(gt_gaze_deg[:, 0])).mean() + |
| torch.abs(y_deg2 - torch.from_numpy(gt_gaze_deg[:, 1])).mean()) / 2 |
| print(f"MAE with idx*2 - 90: {mae2:.4f}") |
| |
| |
| p_err_l = torch.abs(p_deg1 - torch.from_numpy(gt_l_gaze[:, 0] * 180/np.pi)) |
| y_err_l = torch.abs(y_deg1 - torch.from_numpy(gt_l_gaze[:, 1] * 180/np.pi)) |
| mae_l = (p_err_l.mean() + y_err_l.mean()) / 2 |
| |
| p_err_r = torch.abs(p_deg1 - torch.from_numpy(gt_r_gaze[:, 0] * 180/np.pi)) |
| y_err_r = torch.abs(y_deg1 - torch.from_numpy(gt_r_gaze[:, 1] * 180/np.pi)) |
| mae_r = (p_err_r.mean() + y_err_r.mean()) / 2 |
| |
| print(f"Teacher MAE vs Left: {mae_l:.4f}, vs Right: {mae_r:.4f}") |
|
|
| if __name__ == '__main__': |
| check_teacher_accuracy('data/processed/p00.h5') |
| check_teacher_accuracy('data/processed/p01.h5') |
|
|