Gaze-LIPE / src /data /check_teacher_accuracy.py
thanhhuyvan's picture
Publish KD reproducibility investigation
178f61f
Raw
History Blame Contribute Delete
2.32 kB
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'][:]
# Average GT gaze
gt_gaze = (gt_l_gaze + gt_r_gaze) / 2
gt_gaze_deg = gt_gaze * (180.0 / np.pi)
# Convert logits to angles
idx = torch.arange(90).float()
p_prob = F.softmax(p_logits, dim=1)
y_prob = F.softmax(y_logits, dim=1)
# Test 1: Original formula (idx * 4 - 180)
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}")
# Test 2: Formula (idx * 2 - 90)
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}")
# Also check against individual eyes (using Test 1 for now)
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')