Gaze-LIPE / scripts /verify_dann.py
thanhhuyvan's picture
Initial release of LIPE V2 GOLD
a10ba7f
Raw
History Blame Contribute Delete
4.51 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))
from src.models.student import LIPEV2Student
def evaluate_dann_on_gaze360(model_path, h5_path, apply_coord_fix=True):
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
print(f"Evaluating DANN model: {model_path}")
print(f"On dataset: {h5_path}")
print(f"Coordinate Fix (-1, -1): {'ENABLED' if apply_coord_fix else 'DISABLED'}")
print(f"Device: {device}")
# Load Model (Ensure it matches the DANN-enabled architecture)
model = LIPEV2Student().to(device)
state_dict = torch.load(model_path, map_location=device)
model.load_state_dict(state_dict)
model.eval()
results = {
'all': {'error': 0.0, 'count': 0},
'frontal_45': {'error': 0.0, 'count': 0}
}
# Sign multiplier for coordinate alignment
s_p, s_y = (-1, -1) if apply_coord_fix else (1, 1)
with h5py.File(h5_path, 'r') as f:
left_patches = f['left_patches'][:]
right_patches = f['right_patches'][:]
landmarks = f['landmarks'][:]
gaze_gt = f['gaze'][:] # (pitch, yaw) in radians
num_samples = left_patches.shape[0]
with torch.no_grad():
for i in tqdm(range(num_samples), desc="Testing DANN Model"):
# Prepare inputs
lp = torch.from_numpy(left_patches[i]).float().unsqueeze(0).to(device)
rp = torch.from_numpy(right_patches[i]).float().unsqueeze(0).to(device)
lm = torch.from_numpy(landmarks[i]).float().view(1, -1).to(device)
gt = torch.from_numpy(gaze_gt[i]).float().to(device)
# Predict (State A)
# DANN forward returns 3 values: pitch, yaw, domain
p_l, y_l, _ = model(lp, lm, state='A')
p_r, y_r, _ = model(rp, lm, state='A')
# Convert Logits to Degrees
def logits_to_deg(p_logits, y_logits):
idx = torch.arange(90).float().to(device)
p_prob = torch.softmax(p_logits, dim=1)
y_prob = torch.softmax(y_logits, dim=1)
p_deg = (torch.sum(p_prob * idx, dim=1) * 2 - 90)
y_deg = (torch.sum(y_prob * idx, dim=1) * 2 - 90)
return p_deg, y_deg
p_deg_l, y_deg_l = logits_to_deg(p_l, y_l)
p_deg_r, y_deg_r = logits_to_deg(p_r, y_r)
# Average and Apply Coordinate Fix
p_final = ((p_deg_l + p_deg_r) / 2) * s_p
y_final = ((y_deg_l + y_deg_r) / 2) * s_y
# Ground Truth to Degrees
gt_deg = gt * (180.0 / np.pi)
yaw_gt_deg = gt_deg[1].item()
# MAE Calculation
error = (torch.abs(p_final - gt_deg[0]) + torch.abs(y_final - gt_deg[1])).item()
# Update Subsets
results['all']['error'] += error
results['all']['count'] += 1
if abs(yaw_gt_deg) <= 45.0:
results['frontal_45']['error'] += error
results['frontal_45']['count'] += 1
print(f"\n" + "="*45)
print(f"{'SUBSET (DANN ALIGNED)':<20} | {'SAMPLES':<10} | {'MAE (deg)':<10}")
print(f"-"*45)
for key, data in results.items():
if data['count'] > 0:
mae = data['error'] / (data['count'] * 2)
name = "All Cases" if key == 'all' else "Frontal +/- 45"
print(f"{name:<20} | {data['count']:<10} | {mae:.4f}")
print(f"="*45)
if __name__ == "__main__":
# Đánh giá model Experiment 1: DANN Only
model_path = 'checkpoints/dann_only/student_dann_final.pt'
if os.path.exists(model_path):
# Lưu ý: Model DANN Only chưa có AdaLN, nên domain_id sẽ được model tự handle (về 0)
# Hoặc chúng ta có thể sửa script evaluate để linh hoạt hơn.
# Ở đây tôi sẽ pass None cho domain_id vì model hiện tại đã revert về standard LayerNorm.
evaluate_dann_on_gaze360(
model_path=model_path,
h5_path='data/processed/gaze360_robust_v16.h5',
apply_coord_fix=True
)
else:
print(f"Checkpoint not found at {model_path}")