Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
219 changes: 217 additions & 2 deletions robots/libero/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -36,6 +36,8 @@ def get_env_spec() -> EnvSpec:
),
add_cli_args=_add_cli_args,
parse_config=_parse_config,
init_shared_runtime=init_shared_runtime,
init_task_runtime=init_task_runtime,
init_runtime=_init_runtime,
)

Expand All @@ -45,7 +47,6 @@ def get_toolkit(
primitives_kwargs: dict[str, Any],
dashboard_events: DashboardEventSink,
video_path: str | None = None,
save_action_videos: bool = False,
):
"""Return the LIBERO toolkit (common tools + LIBERO primitives)."""
from robots.libero.toolkit import LiberoToolkit
Expand All @@ -54,7 +55,6 @@ def get_toolkit(
primitives_kwargs=primitives_kwargs,
dashboard_events=dashboard_events,
video_path=video_path,
save_action_videos=save_action_videos,
)


Expand Down Expand Up @@ -136,6 +136,212 @@ def _subprocess_env(**extra: str) -> dict[str, str]:
return env


def init_task_runtime(
args: argparse.Namespace,
output_dir: Path,
dashboard_events: DashboardEventSink,
) -> tuple[list[ProcessDaemon], dict[str, Any]]:
"""Initialize one TaskRun-owned LIBERO environment.

A local env server is fresh for every call. When ``--env-endpoint`` is
supplied, the returned daemon list is empty so the external service stays
running.

Heavy runtime dependencies stay lazy so importing :mod:`robots.libero`
for its descriptor or toolkit does not load RPC/model packages.
"""
from robots.libero.env_client import LiberoEnvClient
from rpent.utils.config import get_libero_type
from rpent.utils.daemon import ProcessDaemon, pick_free_port
from rpent.utils.http_rpc import HttpRpcClient
from rpent.utils.rpc import parse_endpoint, wait_for_ready
from rpent.utils.socket_rpc import SocketRpcClient

owned_daemons: list[ProcessDaemon] = []
libero_type = args.libero_type or get_libero_type()
cuda_args = ["--cuda-device", str(args.cuda_device)] if args.cuda_device is not None else []

dashboard_events.emit(RuntimeStatusEvent("env", "starting"))
try:
env_daemon: ProcessDaemon | None = None
if args.env_endpoint is None:
host, port = "127.0.0.1", pick_free_port()
env_daemon = ProcessDaemon(
name="env_server",
cmd=[
sys.executable,
str(get_repo_root() / "robots" / "libero" / "env_server.py"),
"--suite", args.suite,
"--task", str(args.task),
"--seed", str(args.seed),
"--max-episode-steps", str(args.max_episode_steps),
"--transport", "http",
"--host", host,
"--port", str(port),
"--parent-watch",
*cuda_args,
],
env=_subprocess_env(
LIBERO_TYPE=libero_type,
MUJOCO_GL="egl",
ROBOT_PLATFORM="LIBERO",
),
log_path=str(Path(output_dir) / "env_server.log"),
)
env_daemon.start()
owned_daemons.append(env_daemon)
env_rpc: RpcClient = HttpRpcClient(f"http://{host}:{port}")
else:
protocol, host, port = parse_endpoint(args.env_endpoint)
if protocol == "socket":
env_rpc = SocketRpcClient(host, port)
elif protocol == "http":
env_rpc = HttpRpcClient(f"http://{host}:{port}")
else:
raise ValueError(
f"--env-endpoint protocol must be socket or http, got {protocol!r}"
)
wait_for_ready(env_rpc, daemon=env_daemon)
env = LiberoEnvClient(
env_rpc,
expected_meta={
"suite": args.suite,
"task": args.task,
"seed": args.seed,
"max_episode_steps": args.max_episode_steps,
},
)
except Exception as exc:
_stop_owned_daemons(owned_daemons)
dashboard_events.emit(RuntimeStatusEvent("env", "failed", error=exc))
raise
dashboard_events.emit(RuntimeStatusEvent("env", "ready"))
return owned_daemons, {"env": env}


def init_shared_runtime(
args: argparse.Namespace,
output_dir: Path,
dashboard_events: DashboardEventSink,
) -> tuple[list[ProcessDaemon], dict[str, Any]]:
"""Initialize Session-owned VLA and SAM3 services.

The returned list contains only locally started services. External
endpoints are connected to but never become owned.
"""
from rpent.utils.daemon import ProcessDaemon, pick_free_port
from rpent.utils.http_rpc import HttpRpcClient
from rpent.utils.rpc import parse_endpoint, wait_for_ready
from rpent.utils.sam3_client import Sam3Client
from rpent.utils.socket_rpc import SocketRpcClient
from rpent.utils.vla_client import VLAClient

owned_daemons: list[ProcessDaemon] = []
cuda_args = (
["--cuda-device", str(args.cuda_device)]
if args.cuda_device is not None
else []
)

# --- vla_server --------------------------------------------------------
dashboard_events.emit(RuntimeStatusEvent("vla", "starting"))
try:
vla_daemon: ProcessDaemon | None = None
if args.vla_endpoint is None:
host, port = "127.0.0.1", pick_free_port()
vla_daemon = ProcessDaemon(
name="vla_server",
cmd=[
sys.executable,
str(get_repo_root() / "robots" / "libero" / "vla_server.py"),
"--transport", "http",
"--host", host,
"--port", str(port),
"--parent-watch",
*cuda_args,
],
env=_subprocess_env(),
log_path=str(Path(output_dir) / "vla_server.log"),
)
vla_daemon.start()
owned_daemons.append(vla_daemon)
vla_rpc: RpcClient = HttpRpcClient(f"http://{host}:{port}")
else:
protocol, host, port = parse_endpoint(args.vla_endpoint)
if protocol == "socket":
vla_rpc = SocketRpcClient(host, port)
elif protocol == "http":
vla_rpc = HttpRpcClient(f"http://{host}:{port}")
else:
raise ValueError(
f"--vla-endpoint protocol must be socket or http, got {protocol!r}"
)
except Exception as exc:
_stop_owned_daemons(owned_daemons)
dashboard_events.emit(RuntimeStatusEvent("vla", "failed", error=exc))
raise

# --- sam3_server -------------------------------------------------------
dashboard_events.emit(RuntimeStatusEvent("sam3", "starting"))
try:
sam3_daemon: ProcessDaemon | None = None
if args.sam3_endpoint is None:
host, port = "127.0.0.1", pick_free_port()
sam3_daemon = ProcessDaemon(
name="sam3_server",
cmd=[
sys.executable,
str(get_repo_root() / "robots" / "libero" / "sam3_server.py"),
"--transport", "http",
"--host", host,
"--port", str(port),
"--parent-watch",
*cuda_args,
],
env=_subprocess_env(),
log_path=str(Path(output_dir) / "sam3_server.log"),
)
sam3_daemon.start()
owned_daemons.append(sam3_daemon)
sam3_rpc: RpcClient = HttpRpcClient(f"http://{host}:{port}")
else:
protocol, host, port = parse_endpoint(args.sam3_endpoint)
if protocol == "socket":
sam3_rpc = SocketRpcClient(host, port)
elif protocol == "http":
sam3_rpc = HttpRpcClient(f"http://{host}:{port}")
else:
raise ValueError(
f"--sam3-endpoint protocol must be socket or http, got {protocol!r}"
)
except Exception as exc:
_stop_owned_daemons(owned_daemons)
dashboard_events.emit(RuntimeStatusEvent("sam3", "failed", error=exc))
raise

# Start both local services before waiting so heavyweight initialization
# continues concurrently, matching the one-shot runtime behavior.
for component, client, daemon in (
("sam3", sam3_rpc, sam3_daemon),
("vla", vla_rpc, vla_daemon),
):
try:
wait_for_ready(client, daemon=daemon)
except Exception as exc:
_stop_owned_daemons(owned_daemons)
dashboard_events.emit(RuntimeStatusEvent(component, "failed", error=exc))
raise
dashboard_events.emit(RuntimeStatusEvent(component, "ready"))

model = VLAClient(vla_rpc)
sam3_client = Sam3Client(sam3_rpc)

return owned_daemons, {
"model": model,
"sam3_client": sam3_client,
}


def _init_runtime(
args: argparse.Namespace,
output_dir: Path,
Expand Down Expand Up @@ -312,3 +518,12 @@ def _init_runtime(
"sam3_client": Sam3Client(sam3_rpc),
}
return daemons, primitives_kwargs


def _stop_owned_daemons(daemons: list[ProcessDaemon]) -> None:
"""Stop owned daemons in reverse order without masking startup errors."""
for daemon in reversed(daemons):
try:
daemon.stop()
except Exception:
pass
23 changes: 17 additions & 6 deletions robots/libero/toolkit.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,7 @@

from robots.libero import tools as libero_tools
from rpent.dashboard.events import DashboardEventSink, ToolResultEvent
from rpent.tools.toolkit import Toolkit
from rpent.tools.toolkit import ToolCancelled, Toolkit
from rpent.utils.logging import get_logger, get_output_dir


Expand All @@ -29,12 +29,10 @@ def __init__(
primitives_kwargs: dict[str, Any],
dashboard_events: DashboardEventSink,
video_path: str | None = None,
save_action_videos: bool = False,
) -> None:
super().__init__(dashboard_events=dashboard_events)
self._next_step: int = 0
self._video_path: str | None = video_path
self._save_action_videos = save_action_videos
self.init_primitives_clean(primitives_kwargs=primitives_kwargs)
self._register_libero_tools()

Expand Down Expand Up @@ -65,7 +63,15 @@ def _step(self, name: str, **kwargs) -> dict:
command = {"action": name, **kwargs}
t0 = time.time()
start_frame = self._primitives.recorded_frame_count()
result = getattr(self._primitives, name)(**kwargs)
try:
result = getattr(self._primitives, name)(**kwargs)
self.raise_if_cancelled()
except ToolCancelled as exc:
result = {
"error": str(exc),
"code": "tool_cancelled",
"interrupted": True,
}
elapsed = round(time.time() - t0, 2)

if isinstance(result, dict):
Expand All @@ -76,7 +82,7 @@ def _step(self, name: str, **kwargs) -> dict:
self._next_step += 1
step_idx = self._next_step
output_dir = get_output_dir()
if self._save_action_videos:
if self._dashboard_events.enabled:
video_dir = libero_tools.artifact_path(output_dir, "action_videos")
video_path = video_dir / f"step_{step_idx:02d}_{name}.mp4"
try:
Expand All @@ -93,6 +99,8 @@ def _step(self, name: str, **kwargs) -> dict:
)
out = libero_tools.view_driver_state(step_idx)
out["agent_elapsed_s"] = elapsed
if result_dict.get("interrupted"):
out.update(result_dict)
return out

def init_primitives_clean(
Expand All @@ -115,7 +123,10 @@ def init_primitives_clean(
if target.exists():
target.unlink()

primitives = libero_tools.LiberoPrimitives(**primitives_kwargs)
primitives = libero_tools.LiberoPrimitives(
check_cancelled=self.raise_if_cancelled,
**primitives_kwargs,
)
primitives.reset()
primitives.start_recording()
libero_tools.dump_state(primitives, str(out_dir), step_idx=0, log=None)
Expand Down
Loading