Update Wall-X to 1.1.0 (#104)

This commit is contained in:
Starrick Liu
2026-06-15 11:40:00 +08:00
committed by GitHub
parent e23a586846
commit 72834e7de5
200 changed files with 33916 additions and 16771 deletions
+53
View File
@@ -0,0 +1,53 @@
"""Environment registry and factory."""
from __future__ import annotations
from typing import Type
from wall_x._vendor.harrix.eval_config import EvalConfig
from wall_x._vendor.harrix.envs.base import BaseEnv
_REGISTRY: dict[str, Type[BaseEnv]] = {}
def register_env(name: str):
def deco(cls: Type[BaseEnv]):
if name in _REGISTRY:
raise ValueError(
f"env {name!r} already registered (cls={_REGISTRY[name].__name__})"
)
_REGISTRY[name] = cls
return cls
return deco
def _get_class(cfg: EvalConfig) -> Type[BaseEnv]:
t = cfg.env.type
cls = _REGISTRY.get(t)
if cls is None:
if t == "libero":
import wall_x._vendor.harrix.envs as _envs
exc = getattr(_envs, "_LIBERO_IMPORT_ERROR", None)
if exc is not None:
raise RuntimeError(
"LIBERO evaluation dependencies are not installed. "
"Install LIBERO/robosuite and their simulator dependencies "
"before using env.type='libero'."
) from exc
raise ValueError(f"unknown env type={t!r}, registered: {sorted(_REGISTRY)}")
return cls
def build_env(cfg: EvalConfig, worker_id: int) -> BaseEnv:
return _get_class(cfg)(cfg, worker_id)
def enumerate_episodes_for(cfg: EvalConfig) -> list[tuple]:
return _get_class(cfg).enumerate_episodes(cfg)
def registered_envs() -> list[str]:
return sorted(_REGISTRY)