"""Quantify the historical strict=False teacher checkpoint key mismatch.""" from pathlib import Path import json import sys ROOT = Path(r"E:\Gaze_estimation") sys.path.insert(0, str(ROOT / ".codex_deps")) sys.path.insert(0, str(ROOT)) import torch from src.models.teacher import L2CS checkpoint_path = ROOT / "checkpoints" / "resnet50.pt" state = torch.load(checkpoint_path, map_location="cpu", weights_only=True) model = L2CS(backbone_name="resnet50", pretrained=False) result = model.load_state_dict(state, strict=False) model_keys = set(model.state_dict()) checkpoint_keys = set(state) loaded_keys = sorted(model_keys & checkpoint_keys) prefix_map = { "conv1.": "backbone.0.", "bn1.": "backbone.1.", "layer1.": "backbone.4.", "layer2.": "backbone.5.", "layer3.": "backbone.6.", "layer4.": "backbone.7.", } remapped = {} for key, value in state.items(): new_key = key for source, target in prefix_map.items(): if key.startswith(source): new_key = target + key[len(source):] break remapped[new_key] = value fixed_model = L2CS(backbone_name="resnet50", pretrained=False) fixed_result = fixed_model.load_state_dict(remapped, strict=False) report = { "checkpoint": str(checkpoint_path), "checkpoint_tensor_keys": len(checkpoint_keys), "wrapper_tensor_keys": len(model_keys), "historical_loader_loaded_key_count": len(loaded_keys), "historical_loader_loaded_keys": loaded_keys, "historical_loader_missing_key_count": len(result.missing_keys), "historical_loader_unexpected_key_count": len(result.unexpected_keys), "historical_loader_missing_keys": result.missing_keys, "historical_loader_unexpected_keys": result.unexpected_keys, "remapped_loader_missing_keys": fixed_result.missing_keys, "remapped_loader_unexpected_keys": fixed_result.unexpected_keys, "interpretation": "Historical strict=False loading accepts only identically named keys and silently leaves the ResNet backbone at initialization/pretrained values.", } out = ROOT / "artifacts" / "kd-teacher-trap-diagnostic" / "teacher_loader_key_audit.json" out.write_text(json.dumps(report, indent=2), encoding="utf-8") print(json.dumps({k: report[k] for k in ( "checkpoint_tensor_keys", "wrapper_tensor_keys", "historical_loader_loaded_key_count", "historical_loader_loaded_keys", "historical_loader_missing_key_count", "historical_loader_unexpected_key_count", "remapped_loader_missing_keys", "remapped_loader_unexpected_keys")}, indent=2))