from pathlib import Path from PIL import Image from spelunkai_training.client import DatasetVersionData, FrameData, LabelData from spelunkai_training.dataset import CenterNetDataset from spelunkai_training.model import OUTPUT_STRIDE def _write_fake_image(path: Path, width: int, height: int) -> None: path.parent.mkdir(parents=True, exist_ok=True) Image.new("RGB", (width, height), color=(10, 20, 30)).save(path) def test_dataset_getitem_shapes(tmp_path): images_root = tmp_path / "frames" _write_fake_image(images_root / "run01_frames" / "frame_000000.png", width=64, height=64) dataset_version = DatasetVersionData( id=1, set_id=1, name="v1", frames=[ FrameData( id=1, session_name="run01", frame_index=0, image_path="run01_frames/frame_000000.png", width=64, height=64, labels=[LabelData(sub_class_id=7, x=8, y=8, width=16, height=16)], ), ], ) dataset = CenterNetDataset( dataset_version=dataset_version, images_root=images_root, class_index_by_sub_class_id={7: 0}, ) assert len(dataset) == 1 sample = dataset[0] out_size = 64 // OUTPUT_STRIDE assert sample["image"].shape == (3, 64, 64) assert sample["heatmap"].shape == (1, out_size, out_size) assert sample["wh"].shape == (2, out_size, out_size) assert sample["mask"].shape == (out_size, out_size) assert sample["mask"].sum().item() == 1.0 def test_dataset_ignores_labels_for_unknown_sub_classes(tmp_path): images_root = tmp_path / "frames" _write_fake_image(images_root / "run01_frames" / "frame_000000.png", width=32, height=32) dataset_version = DatasetVersionData( id=1, set_id=1, name="v1", frames=[ FrameData( id=1, session_name="run01", frame_index=0, image_path="run01_frames/frame_000000.png", width=32, height=32, labels=[LabelData(sub_class_id=999, x=0, y=0, width=4, height=4)], ), ], ) dataset = CenterNetDataset( dataset_version=dataset_version, images_root=images_root, class_index_by_sub_class_id={7: 0}, ) sample = dataset[0] assert sample["mask"].sum().item() == 0.0