feat: model registry scans models/ directory with known class map
This commit is contained in:
1 parent
71c1ae1cf2
commit
f5982a4222
2 files changed
+101
No files matched your search
@@ -0,0 +1,54 @@
|
||||
"""Model registry — scans models/ directory and returns available model configs."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
KNOWN_MODEL_CLASSES: dict[str, list[str]] = {
|
||||
"truck-detector": ["truck"],
|
||||
"v4-best": ["sack", "truck"],
|
||||
"model_karung_truk": ["sack", "truck"],
|
||||
"karung-dimuat-detection-di-feedmill-yolo26n-seg-200e": ["person", "sack"],
|
||||
"yolo11n-bbox-100ep-sack+box-20260909-best": ["sack", "box"],
|
||||
"best": ["sack"],
|
||||
}
|
||||
|
||||
MODEL_EXTENSIONS = {".pt", ".onnx", ".engine"}
|
||||
|
||||
|
||||
@dataclass
|
||||
class ModelConfig:
|
||||
"""A discovered model weight file with metadata."""
|
||||
|
||||
filename: str
|
||||
path: str
|
||||
stem: str
|
||||
known_classes: list[str] = field(default_factory=list)
|
||||
|
||||
|
||||
def scan_models(models_dir: str) -> list[ModelConfig]:
|
||||
"""Scan models_dir for weight files and return ModelConfig list.
|
||||
|
||||
Sorts by filename for stable ordering.
|
||||
"""
|
||||
p = Path(models_dir)
|
||||
if not p.is_dir():
|
||||
return []
|
||||
|
||||
configs: list[ModelConfig] = []
|
||||
for f in sorted(p.iterdir()):
|
||||
if f.is_file() and f.suffix in MODEL_EXTENSIONS:
|
||||
stem = f.stem
|
||||
known = KNOWN_MODEL_CLASSES.get(stem, [])
|
||||
configs.append(
|
||||
ModelConfig(
|
||||
filename=f.name,
|
||||
path=str(f.resolve()),
|
||||
stem=stem,
|
||||
known_classes=list(known),
|
||||
)
|
||||
)
|
||||
return configs
|
||||
@@ -0,0 +1,47 @@
|
||||
"""Tests for model registry (src/model_registry.py)."""
|
||||
|
||||
import pytest
|
||||
from src.model_registry import scan_models, ModelConfig
|
||||
|
||||
|
||||
def test_scan_returns_list():
|
||||
result = scan_models("/nonexistent/path")
|
||||
assert isinstance(result, list)
|
||||
|
||||
|
||||
def test_scan_empty_dir(tmp_path):
|
||||
result = scan_models(str(tmp_path))
|
||||
assert result == []
|
||||
|
||||
|
||||
def test_scan_finds_pt_files(tmp_path):
|
||||
(tmp_path / "best.pt").write_bytes(b"fake")
|
||||
(tmp_path / "truck-detector.pt").write_bytes(b"fake")
|
||||
result = scan_models(str(tmp_path))
|
||||
assert len(result) == 2
|
||||
names = {m.filename for m in result}
|
||||
assert "best.pt" in names
|
||||
assert "truck-detector.pt" in names
|
||||
|
||||
|
||||
def test_scan_skips_non_model_files(tmp_path):
|
||||
(tmp_path / "modelREADME.md").write_text("readme")
|
||||
(tmp_path / "best.pt").write_bytes(b"fake")
|
||||
result = scan_models(str(tmp_path))
|
||||
assert len(result) == 1
|
||||
|
||||
|
||||
def test_model_config_fields(tmp_path):
|
||||
(tmp_path / "v4-best.pt").write_bytes(b"fake")
|
||||
result = scan_models(str(tmp_path))
|
||||
cfg = result[0]
|
||||
assert cfg.filename == "v4-best.pt"
|
||||
assert cfg.path == str(tmp_path / "v4-best.pt")
|
||||
assert isinstance(cfg.known_classes, list)
|
||||
|
||||
|
||||
def test_model_config_fallback_classes(tmp_path):
|
||||
(tmp_path / "unknown-model.pt").write_bytes(b"fake")
|
||||
result = scan_models(str(tmp_path))
|
||||
cfg = result[0]
|
||||
assert cfg.known_classes == []
|
||||
Reference in new issue
Block a user