feat: add ModelGroup and scan_model_groups() for grouped model display
This commit is contained in:
1 parent
d32045b433
commit
22c78429cf
2 files changed
+162
-1
No files matched your search
@@ -28,6 +28,26 @@ class ModelConfig:
|
|||||||
known_classes: list[str] = field(default_factory=list)
|
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]:
|
def scan_models(models_dir: str) -> list[ModelConfig]:
|
||||||
"""Scan models_dir for weight files and return ModelConfig list.
|
"""Scan models_dir for weight files and return ModelConfig list.
|
||||||
|
|
||||||
@@ -51,3 +71,36 @@ def scan_models(models_dir: str) -> list[ModelConfig]:
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
return configs
|
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
|
||||||
@@ -1,7 +1,14 @@
|
|||||||
"""Tests for model registry (src/model_registry.py)."""
|
"""Tests for model registry (src/model_registry.py)."""
|
||||||
|
|
||||||
import pytest
|
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():
|
def test_scan_returns_list():
|
||||||
@@ -45,3 +52,104 @@ def test_model_config_fallback_classes(tmp_path):
|
|||||||
result = scan_models(str(tmp_path))
|
result = scan_models(str(tmp_path))
|
||||||
cfg = result[0]
|
cfg = result[0]
|
||||||
assert cfg.known_classes == []
|
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"]
|
||||||
Reference in new issue
Block a user