feat: model registry scans models/ directory with known class map

This commit is contained in:
jetson committed 2026-09-16 14:50:35 +07:00
1 parent 71c1ae1cf2
commit f5982a4222
2 files changed
+101

No files matched your search

+54
View File
@@ -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
+47
View File
@@ -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 == []