| import torch |
| import sys |
| import os |
| from pathlib import Path |
|
|
| |
| sys.path.append(str(Path(__file__).parent.parent.parent)) |
| from src.models.teacher import load_teacher_model |
|
|
| def verify_teacher(): |
| checkpoint_path = 'checkpoints/resnet50.pt' |
| if not os.path.exists(checkpoint_path): |
| print(f"Error: {checkpoint_path} not found.") |
| return False |
| |
| print(f"Loading teacher model from {checkpoint_path}...") |
| try: |
| model = load_teacher_model(checkpoint_path, backbone='resnet50', device='cpu') |
| print("Teacher model loaded successfully!") |
| |
| |
| for name, param in model.named_parameters(): |
| if param.dtype != torch.float32: |
| print(f"Warning: Parameter {name} is {param.dtype}") |
| break |
| |
| |
| dummy_input = torch.randn(1, 3, 224, 224) |
| with torch.no_grad(): |
| p_logits, y_logits = model(dummy_input) |
| print(f"Output shapes: Pitch {p_logits.shape}, Yaw {y_logits.shape}") |
| |
| p_deg, y_deg = model.get_angles(p_logits, y_logits) |
| print(f"Predicted angles (dummy): Pitch {p_deg.item():.2f}, Yaw {y_deg.item():.2f}") |
| |
| return True |
| except Exception as e: |
| print(f"Error loading teacher: {e}") |
| |
| checkpoint = torch.load(checkpoint_path, map_location='cpu') |
| print(f"Keys in checkpoint: {list(checkpoint.keys())[:10]}...") |
| return False |
|
|
| if __name__ == '__main__': |
| verify_teacher() |
|
|