pi0.5 / model_loader.py
XiangpengYang's picture
feat: select UR policy configuration
380cc0c
Raw
History Blame Contribute Delete
3.68 kB
"""Thread-safe lazy lifecycle for the heavyweight π₀.₅ UR policy."""
from __future__ import annotations
import gc
from pathlib import Path
import sys
import threading
from collections.abc import Callable
from artifacts import (
download_checkpoint,
normalize_checkpoint_path,
normalize_model_id,
)
POLICY_CONFIGS = ("pi05_ur_demo_no_state", "pi05_ur_demo_state")
DEFAULT_POLICY_CONFIG = POLICY_CONFIGS[0]
def normalize_policy_config(value: str) -> str:
if value not in POLICY_CONFIGS:
choices = ", ".join(POLICY_CONFIGS)
raise ValueError(f"unsupported policy config {value!r}; choose one of: {choices}")
return value
class ModelUnavailableError(RuntimeError):
"""Raised when the requested policy cannot be initialized."""
def _release_gpu_memory() -> None:
gc.collect()
try:
import torch
if torch.cuda.is_available():
torch.cuda.empty_cache()
except ImportError:
pass
class ModelManager:
def __init__(self, loader: Callable[[str, str, str], object] | None = None):
self._loader = loader or self._load_default
self._lock = threading.Lock()
self._value = None
self._active_key: tuple[str, str, str] | None = None
self._error: str | None = None
@property
def active_key(self) -> tuple[str, str, str] | None:
return self._active_key
@property
def health_message(self) -> str:
if self._value is not None and self._active_key is not None:
model_id, checkpoint_path, config_name = self._active_key
return f"Model ready: {model_id}/{checkpoint_path} ({config_name})."
if self._error:
return f"Model unavailable: {self._error}"
return "Model has not been loaded yet."
def get(self, model_id: str, checkpoint_path: str, config_name: str):
key = (
normalize_model_id(model_id),
normalize_checkpoint_path(checkpoint_path),
normalize_policy_config(config_name),
)
if self._value is not None and self._active_key == key:
return self._value
with self._lock:
if self._value is not None and self._active_key == key:
return self._value
if self._value is not None:
self._value = None
self._active_key = None
_release_gpu_memory()
self._error = None
try:
value = self._loader(*key)
except Exception as exc:
_release_gpu_memory()
detail = str(exc) or exc.__class__.__name__
self._error = f"{key[0]}/{key[1]} ({key[2]}): {detail}"
raise ModelUnavailableError(self._error) from exc
self._value = value
self._active_key = key
return value
@staticmethod
def _load_default(model_id: str, checkpoint_path: str, config_name: str):
import torch
if not torch.cuda.is_available():
raise RuntimeError("CUDA GPU is required for π₀.₅ inference")
runtime = str(Path(__file__).resolve().parent / "openpi_runtime")
if runtime not in sys.path:
sys.path.insert(0, runtime)
from openpi.policies import policy_config
from openpi.training import config as openpi_config
paths = download_checkpoint(model_id, checkpoint_path)
config = openpi_config.get_config(config_name)
return policy_config.create_trained_policy(
config,
paths.checkpoint,
pytorch_device="cuda",
)
MODEL_MANAGER = ModelManager()