fabric-defect-api / model_loader.py
rafatkabir's picture
Create model_loader.py
3c59252 verified
Raw
History Blame Contribute Delete
2.35 kB
import os
import json
import logging
from pathlib import Path
from huggingface_hub import hf_hub_download
logger = logging.getLogger(__name__)
# Injected from Space secrets
MODEL_REPO = os.environ.get("HF_MODEL_REPO", "")
HF_TOKEN = os.environ.get("HF_TOKEN", "")
# /tmp is the only writable directory in HF Spaces
MODEL_CACHE = Path("/tmp/model_cache")
ONNX_PATH = MODEL_CACHE / "best.onnx"
CLASS_MAP_PATH = MODEL_CACHE / "class_map.json"
_model = None
_class_names = {}
_model_version = "unknown"
def load_model():
global _model, _class_names, _model_version
MODEL_CACHE.mkdir(parents=True, exist_ok=True)
# Download model weights if not already in /tmp cache
if not ONNX_PATH.exists():
if not MODEL_REPO:
raise RuntimeError(
"HF_MODEL_REPO is not set. "
"Add it under Space Settings → Repository secrets."
)
logger.info(f"Downloading best.onnx from {MODEL_REPO} ...")
hf_hub_download(
repo_id = MODEL_REPO,
filename = "best.onnx",
local_dir = str(MODEL_CACHE),
token = HF_TOKEN or None,
)
logger.info("Download complete.")
else:
logger.info("Model found in /tmp cache — skipping download.")
# Download class map if missing
if not CLASS_MAP_PATH.exists():
logger.info("Downloading class_map.json ...")
hf_hub_download(
repo_id = MODEL_REPO,
filename = "class_map.json",
local_dir = str(MODEL_CACHE),
token = HF_TOKEN or None,
)
# Load ONNX model via ultralytics
from ultralytics import YOLO
_model = YOLO(str(ONNX_PATH))
logger.info("YOLO model loaded.")
# Load class map {"0": "no_defect", "1": "hole", ...}
with open(CLASS_MAP_PATH) as f:
raw = json.load(f)
_class_names = {int(k): v for k, v in raw.items()}
_model_version = MODEL_REPO.split("/")[-1] if MODEL_REPO else "local"
logger.info(f"Ready. {len(_class_names)} classes: {list(_class_names.values())}")
def get_model():
if _model is None:
raise RuntimeError("Model not loaded yet.")
return _model
def get_class_names() -> dict:
return _class_names
def get_model_version() -> str:
return _model_version