diff --git a/src/model_registry.py b/src/model_registry.py index 0f6f60b..5b75965 100644 --- a/src/model_registry.py +++ b/src/model_registry.py @@ -28,6 +28,26 @@ class ModelConfig: known_classes: list[str] = field(default_factory=list) +FORMAT_PREFERENCE: list[str] = [".engine", ".pt", ".onnx"] + + +@dataclass +class ModelGroup: + """A group of model files sharing the same stem, with format options.""" + + stem: str + formats: list[str] + format_paths: dict[str, str] + known_classes: list[str] + default_format: str + + +def _sort_formats(formats: list[str]) -> list[str]: + """Sort formats by preference: .engine > .pt > .onnx, unknowns last.""" + order = {ext: i for i, ext in enumerate(FORMAT_PREFERENCE)} + return sorted(formats, key=lambda f: order.get(f, len(FORMAT_PREFERENCE))) + + def scan_models(models_dir: str) -> list[ModelConfig]: """Scan models_dir for weight files and return ModelConfig list. @@ -51,3 +71,36 @@ def scan_models(models_dir: str) -> list[ModelConfig]: ) ) return configs + + +def scan_model_groups(models_dir: str) -> list[ModelGroup]: + """Scan models_dir and group discovered files by stem. + + Returns a list of ModelGroup sorted by stem for stable ordering. + """ + p = Path(models_dir) + if not p.is_dir(): + return [] + + stem_map: dict[str, dict[str, str]] = {} + for f in sorted(p.iterdir()): + if f.is_file() and f.suffix.lower() in MODEL_EXTENSIONS: + stem = f.stem + ext = f.suffix.lower() + if stem not in stem_map: + stem_map[stem] = {} + stem_map[stem][ext] = str(f.resolve()) + + groups: list[ModelGroup] = [] + for stem in sorted(stem_map): + formats = _sort_formats(list(stem_map[stem].keys())) + groups.append( + ModelGroup( + stem=stem, + formats=formats, + format_paths=stem_map[stem], + known_classes=list(KNOWN_MODEL_CLASSES.get(stem, [])), + default_format=formats[0], + ) + ) + return groups diff --git a/tests/test_model_registry.py b/tests/test_model_registry.py index 4ba6893..db844fb 100644 --- a/tests/test_model_registry.py +++ b/tests/test_model_registry.py @@ -1,7 +1,14 @@ """Tests for model registry (src/model_registry.py).""" import pytest -from src.model_registry import scan_models, ModelConfig +from src.model_registry import ( + scan_models, + scan_model_groups, + ModelConfig, + ModelGroup, + _sort_formats, + FORMAT_PREFERENCE, +) def test_scan_returns_list(): @@ -45,3 +52,104 @@ def test_model_config_fallback_classes(tmp_path): result = scan_models(str(tmp_path)) cfg = result[0] assert cfg.known_classes == [] + + +# --- scan_model_groups tests --- + + +def test_scan_model_groups_returns_list(): + result = scan_model_groups("/nonexistent/path") + assert isinstance(result, list) + + +def test_scan_model_groups_empty_dir(tmp_path): + result = scan_model_groups(str(tmp_path)) + assert result == [] + + +def test_scan_model_groups_single_stem(tmp_path): + (tmp_path / "best.pt").write_bytes(b"fake") + (tmp_path / "best.onnx").write_bytes(b"fake") + (tmp_path / "best.engine").write_bytes(b"fake") + result = scan_model_groups(str(tmp_path)) + assert len(result) == 1 + g = result[0] + assert g.stem == "best" + assert ".engine" in g.formats + assert ".pt" in g.formats + assert ".onnx" in g.formats + assert g.default_format == ".engine" + + +def test_scan_model_groups_format_preference(tmp_path): + (tmp_path / "model.pt").write_bytes(b"fake") + (tmp_path / "model.onnx").write_bytes(b"fake") + result = scan_model_groups(str(tmp_path)) + g = result[0] + assert g.formats == [".pt", ".onnx"] + assert g.default_format == ".pt" + + +def test_scan_model_groups_sorted_by_stem(tmp_path): + (tmp_path / "z-model.pt").write_bytes(b"fake") + (tmp_path / "a-model.pt").write_bytes(b"fake") + (tmp_path / "m-model.pt").write_bytes(b"fake") + result = scan_model_groups(str(tmp_path)) + stems = [g.stem for g in result] + assert stems == sorted(stems) + + +def test_scan_model_groups_known_classes(tmp_path): + (tmp_path / "truck-detector.pt").write_bytes(b"fake") + (tmp_path / "truck-detector.onnx").write_bytes(b"fake") + result = scan_model_groups(str(tmp_path)) + g = result[0] + assert g.known_classes == ["truck"] + + +def test_scan_model_groups_unknown_stem(tmp_path): + (tmp_path / "unknown-model.pt").write_bytes(b"fake") + result = scan_model_groups(str(tmp_path)) + g = result[0] + assert g.known_classes == [] + + +def test_scan_model_groups_format_paths(tmp_path): + (tmp_path / "best.pt").write_bytes(b"fake") + (tmp_path / "best.engine").write_bytes(b"fake") + result = scan_model_groups(str(tmp_path)) + g = result[0] + assert ".pt" in g.format_paths + assert ".engine" in g.format_paths + assert g.format_paths[".pt"] == str((tmp_path / "best.pt").resolve()) + assert g.format_paths[".engine"] == str((tmp_path / "best.engine").resolve()) + + +def test_scan_model_groups_skips_non_model_files(tmp_path): + (tmp_path / "best.pt").write_bytes(b"fake") + (tmp_path / "readme.md").write_text("hi") + result = scan_model_groups(str(tmp_path)) + assert len(result) == 1 + assert result[0].stem == "best" + + +def test_scan_model_groups_multiple_stems(tmp_path): + (tmp_path / "best.pt").write_bytes(b"fake") + (tmp_path / "v4-best.engine").write_bytes(b"fake") + (tmp_path / "truck-detector.onnx").write_bytes(b"fake") + result = scan_model_groups(str(tmp_path)) + assert len(result) == 3 + stems = [g.stem for g in result] + assert "best" in stems + assert "v4-best" in stems + assert "truck-detector" in stems + + +def test_sort_formats_explicit(): + result = _sort_formats([".onnx", ".pt", ".engine"]) + assert result == [".engine", ".pt", ".onnx"] + + +def test_sort_formats_unknown_ext(): + result = _sort_formats([".pt", ".xyz", ".onnx"]) + assert result == [".pt", ".onnx", ".xyz"]