ocr-services / src /storage.py
Al-Fathir
fix : load model
5d06e53
Raw
History Blame Contribute Delete
2.21 kB
"""
Centralized path management for persistent model storage.
HF Spaces mounts persistent storage at /data. Local development falls back to
~/.cache/ocr-demo so model downloads survive app restarts.
"""
from __future__ import annotations
import os
from pathlib import Path
_IS_HF_SPACE = os.path.isdir("/data") or os.environ.get("SPACE_ID") is not None
if _IS_HF_SPACE:
_BASE = Path("/data")
else:
_BASE = Path(os.getenv("OCR_DEMO_CACHE_DIR", Path.home() / ".cache" / "ocr-demo"))
PATHS = {
"hf_home": _BASE / "huggingface",
"tmp": Path("/tmp") / "ocr-demo",
}
def ensure_dirs() -> None:
for path in PATHS.values():
path.mkdir(parents=True, exist_ok=True)
def get_env_overrides() -> dict[str, str]:
hf_home = PATHS["hf_home"]
return {
"HF_HOME": str(hf_home),
"HUGGINGFACE_HUB_CACHE": str(hf_home / "hub"),
"TRANSFORMERS_CACHE": str(hf_home / "hub"),
"HF_XET_HIGH_PERFORMANCE": os.getenv("HF_XET_HIGH_PERFORMANCE", "1"),
"CUDA_VISIBLE_DEVICES": os.getenv("CUDA_VISIBLE_DEVICES", ""),
"PYTHONUNBUFFERED": "1",
}
def is_model_cached(model_id: str) -> bool:
repo_dir = f"models--{model_id.replace('/', '--')}"
snapshots = PATHS["hf_home"] / "hub" / repo_dir / "snapshots"
if not snapshots.exists():
return False
weight_patterns = (
"model*.safetensors",
"*.safetensors",
"pytorch_model*.bin",
"model*.bin",
"*.gguf",
)
for snapshot in snapshots.iterdir():
if not snapshot.is_dir():
continue
for pattern in weight_patterns:
if any(snapshot.glob(pattern)):
return True
return False
def storage_info() -> dict:
import shutil
info = {
"environment": "HuggingFace Spaces" if _IS_HF_SPACE else "Local Dev",
"base_path": str(_BASE),
"paths": {k: str(v) for k, v in PATHS.items()},
}
if _BASE.exists():
total, used, free = shutil.disk_usage(str(_BASE))
info["disk"] = {
"total_gb": round(total / 1e9, 2),
"used_gb": round(used / 1e9, 2),
"free_gb": round(free / 1e9, 2),
}
return info