import torch from spelunkai_training.model import OUTPUT_STRIDE, CenterNetDetector def test_forward_pass_output_shapes(): model = CenterNetDetector(num_classes=3, base_channels=8) model.eval() x = torch.rand(2, 3, 64, 64) with torch.no_grad(): heatmap, wh = model(x) expected_size = 64 // OUTPUT_STRIDE assert heatmap.shape == (2, 3, expected_size, expected_size) assert wh.shape == (2, 2, expected_size, expected_size) def test_heatmap_output_is_in_unit_range(): model = CenterNetDetector(num_classes=2, base_channels=8) model.eval() x = torch.rand(1, 3, 32, 32) with torch.no_grad(): heatmap, _ = model(x) assert heatmap.min() >= 0.0 assert heatmap.max() <= 1.0 def test_real_capture_resolution_is_divisible_by_output_stride(): # Spelunky Classic HD's fixed capture resolution (CLAUDE.md ยง3.1). assert 1280 % OUTPUT_STRIDE == 0 assert 720 % OUTPUT_STRIDE == 0