Spaces:
Runtime error
Runtime error
Download app.py from zaleni/wsa1-robot-control: direct link, hf CLI and curl.
- Browser
- Download file 18.1 kB
-
https://huggingface.co/spaces/zaleni/wsa1-robot-control/resolve/main/app.py
- Command line
-
hf download hf://spaces/zaleni/wsa1-robot-control/app.py
-
curl -L -o app.py https://huggingface.co/spaces/zaleni/wsa1-robot-control/resolve/main/app.py
18.1 kB
| """WSA-Base — 3D-Centric World-Spatial-Action model for generalizable robot control. | |
| Interactive demo: given two camera views (head + wrist) of a robot scene and a | |
| natural-language instruction, the WSA-Base policy predicts a short chunk of | |
| end-effector actions AND (via its action-conditioned world model) a decoded | |
| prediction of the next head-camera frame. | |
| Model: zaleni/WSA-Base-LIBERO (3B, Qwen3-VL-2B backbone, LIBERO franka adapter). | |
| """ | |
| import os | |
| os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True") | |
| os.environ.setdefault("TOKENIZERS_PARALLELISM", "false") | |
| import sys | |
| import shutil | |
| from pathlib import Path | |
| APP_DIR = Path(__file__).resolve().parent | |
| # --- Patch transformers' Qwen3-VL modeling with WSA's cached-inference variant --- | |
| # WSA ships a small patch to transformers.models.qwen3_vl.modeling_qwen3_vl that | |
| # adds a use_cache=False KV-concat path used by its multi-expert cached inference. | |
| # The upstream repo installs it by copying the file over site-packages; we do the | |
| # same here BEFORE importing transformers, so the patched module is the one loaded. | |
| def _patch_transformers_qwen3_vl(): | |
| try: | |
| import transformers | |
| except Exception as e: # pragma: no cover | |
| print(f"[patch] transformers import failed: {e!r}") | |
| return | |
| patched = ( | |
| APP_DIR | |
| / "lerobot" | |
| / "policies" | |
| / "WSA_Base" | |
| / "transformers_replace" | |
| / "models" | |
| / "qwen3_vl" | |
| / "modeling_qwen3_vl.py" | |
| ) | |
| dest = Path(transformers.__file__).parent / "models" / "qwen3_vl" / "modeling_qwen3_vl.py" | |
| if not patched.exists(): | |
| print(f"[patch] patched file missing at {patched}") | |
| return | |
| try: | |
| if dest.exists() and dest.read_text() == patched.read_text(): | |
| print("[patch] qwen3_vl already patched") | |
| return | |
| shutil.copy2(patched, dest) | |
| print(f"[patch] copied WSA qwen3_vl modeling -> {dest}") | |
| except Exception as e: # pragma: no cover | |
| print(f"[patch] failed to patch qwen3_vl modeling: {e!r}") | |
| _patch_transformers_qwen3_vl() | |
| # lerobot is bundled in ./lerobot ; make it importable. | |
| if str(APP_DIR) not in sys.path: | |
| sys.path.insert(0, str(APP_DIR)) | |
| import spaces # noqa: E402 (must precede torch / CUDA-touching imports) | |
| import json # noqa: E402 | |
| import time # noqa: E402 | |
| import tempfile # noqa: E402 | |
| import numpy as np # noqa: E402 | |
| import torch # noqa: E402 | |
| import gradio as gr # noqa: E402 | |
| from PIL import Image # noqa: E402 | |
| import matplotlib # noqa: E402 | |
| matplotlib.use("Agg") | |
| import matplotlib.pyplot as plt # noqa: E402 | |
| from huggingface_hub import snapshot_download # noqa: E402 | |
| # WSA / lerobot imports | |
| from lerobot.configs.policies import PreTrainedConfig # noqa: E402 | |
| from lerobot.policies.WSA_Base.configuration_wsa_base import WSABaseConfig # noqa: E402, F401 (registers "TBot_SA1") | |
| from lerobot.policies.WSA_Base.modeling_wsa_base import WSABasePolicy # noqa: E402 | |
| from lerobot.policies.WSA_Base.transform_wsa_base import ( # noqa: E402 | |
| Qwen3_VLProcessorTransformFn as WSABaseProcessorTransformFn, | |
| ) | |
| from lerobot.transforms.core import ( # noqa: E402 | |
| NormalizeTransformFn, | |
| ResizeImagesWithPadFn, | |
| UnNormalizeTransformFn, | |
| ) | |
| from lerobot.utils.constants import OBS_IMAGES, OBS_STATE # noqa: E402 | |
| MODEL_ID = "zaleni/WSA-Base-LIBERO" | |
| RESIZE = 224 | |
| DTYPE = torch.bfloat16 | |
| # LIBERO 7-DoF absolute end-effector control (action_mode="abs"): | |
| # [x, y, z, roll, pitch, yaw, gripper] | |
| ACTION_LABELS = ["x", "y", "z", "roll", "pitch", "yaw", "gripper"] | |
| # Number of raw action dims for LIBERO (stats are 7-dim; model pads to 32). | |
| TARGET_ACTION_DIM = 7 | |
| # Neutral LIBERO franka start state (8-dim: eef_pos[3] + eef_axisangle[3] + gripper[2]). | |
| DEFAULT_STATE = [-0.203, 0.010, 1.178, 3.140, 0.004, -0.093, 0.039, -0.039] | |
| # --------------------------------------------------------------------------- | |
| # Model load (module scope, eager to CUDA — ZeroGPU packs weights on startup). | |
| # --------------------------------------------------------------------------- | |
| print("[startup] resolving checkpoint ...") | |
| _ckpt_dir = Path(snapshot_download(repo_id=MODEL_ID)) | |
| _config = PreTrainedConfig.from_pretrained(_ckpt_dir) | |
| # On ZeroGPU there is no CUDA device at startup — build everything on CPU and | |
| # move to the GPU lazily inside the @spaces.GPU function. The cosmos tokenizer | |
| # loads a TorchScript module directly onto `config.(cosmos_)device`, so this | |
| # MUST be "cpu" at construction time or startup crashes with cudaErrorNoDevice. | |
| _config.device = "cpu" | |
| if hasattr(_config, "cosmos_device"): | |
| _config.cosmos_device = "cpu" | |
| # Disable the DA3 3D teacher for eval (matches DISABLE_DA3_TEACHER_FOR_EVAL=true). | |
| if hasattr(_config, "lambda_3d"): | |
| _config.lambda_3d = 0.0 | |
| print("[startup] loading WSA-Base policy (CPU) ...") | |
| policy = WSABasePolicy.from_pretrained(config=_config, pretrained_name_or_path=_ckpt_dir) | |
| policy = policy.to(dtype=DTYPE).eval() | |
| print("[startup] policy loaded on CPU.") | |
| _ON_CUDA = False | |
| def _ensure_cuda(): | |
| """Move the policy (incl. the JIT cosmos tokenizer) to CUDA once, lazily.""" | |
| global _ON_CUDA | |
| if _ON_CUDA: | |
| return | |
| policy.to(device="cuda", dtype=DTYPE) | |
| # The cosmos ImageTokenizer wraps TorchScript encoder/decoder modules that | |
| # were loaded as float32 on CPU. Move them to CUDA and cast to bf16 so the | |
| # world-model latents (bf16) match, mirroring the CUDA construction path. | |
| try: | |
| cosmos = policy.model.cosmos | |
| for attr in ("_full_model", "_enc_model", "_dec_model"): | |
| m = getattr(cosmos, attr, None) | |
| if m is not None and hasattr(m, "to"): | |
| m.to(device="cuda", dtype=DTYPE) | |
| cosmos._device = "cuda" | |
| cosmos._dtype = DTYPE | |
| except Exception as exc: # noqa: BLE001 | |
| print(f"[warn] cosmos .to(cuda) partial: {exc}") | |
| _ON_CUDA = True | |
| # Stats for state/action (un)normalization. | |
| with open(_ckpt_dir / "stats.json") as f: | |
| _stats_root = json.load(f) | |
| _stats_key = next(iter(_stats_root)) if len(_stats_root) == 1 else _stats_root.get("franka", None) | |
| _stats = _stats_root[_stats_key] if isinstance(_stats_key, str) else list(_stats_root.values())[0] | |
| _STAT_KEYS = ["min", "max", "mean", "std"] | |
| _state_stat = {"observation.state": {k: np.asarray(_stats["observation.state"][k]) for k in _STAT_KEYS}} | |
| _action_stat = {"action": {k: np.asarray(_stats["action"][k]) for k in _STAT_KEYS}} | |
| _resize_fn = ResizeImagesWithPadFn(height=RESIZE, width=RESIZE) | |
| _normalize_state_fn = NormalizeTransformFn( | |
| selected_keys=["observation.state"], mode="mean_std", norm_stats=_state_stat | |
| ) | |
| _unnormalize_action_fn = UnNormalizeTransformFn( | |
| selected_keys=["action"], mode="mean_std", norm_stats=_action_stat | |
| ) | |
| _processor_path = ( | |
| getattr(_config, "qwen3_vl_processor_path", None) | |
| or getattr(_config, "qwen3_vl_pretrained_path", None) | |
| or "Qwen/Qwen3-VL-2B-Instruct" | |
| ) | |
| _processor_fn = WSABaseProcessorTransformFn( | |
| pretrained_model_name_or_path=_processor_path, | |
| max_length=int(getattr(_config, "tokenizer_max_length", 48)), | |
| ) | |
| _ACTION_DIM = _config.output_features["action"].shape[0] | |
| _CHUNK = int(getattr(_config, "chunk_size", 10)) | |
| # --------------------------------------------------------------------------- | |
| # Preprocessing helpers (mirrors evaluation/Libero/inference.py). | |
| # --------------------------------------------------------------------------- | |
| def _to_history_tensor(img: Image.Image) -> torch.Tensor: | |
| """PIL RGB image -> [T=2, H, W, C] float tensor in [0,1] (duplicated frame).""" | |
| arr = np.asarray(img.convert("RGB"), dtype=np.float32) / 255.0 | |
| frame = torch.from_numpy(arr) | |
| return torch.stack([frame, frame], dim=0) # T=2 (past, current) | |
| def _prepare_inputs(head_img, wrist_img, state, instruction): | |
| head_hist = _to_history_tensor(head_img) | |
| wrist_hist = _to_history_tensor(wrist_img) | |
| dummy_hist = torch.ones_like(head_hist) | |
| sample = { | |
| f"{OBS_IMAGES}.image0": head_hist.permute(0, 3, 1, 2), | |
| f"{OBS_IMAGES}.image1": wrist_hist.permute(0, 3, 1, 2), | |
| f"{OBS_IMAGES}.image2": dummy_hist.permute(0, 3, 1, 2), | |
| OBS_STATE: torch.from_numpy(np.asarray(state, dtype=np.float32)), | |
| "task": instruction, | |
| } | |
| sample = _resize_fn(sample) | |
| sample[f"{OBS_IMAGES}.image0_mask"] = torch.tensor(True) | |
| sample[f"{OBS_IMAGES}.image1_mask"] = torch.tensor(True) | |
| sample[f"{OBS_IMAGES}.image2_mask"] = torch.tensor(False) | |
| sample = _processor_fn(sample) | |
| sample = _normalize_state_fn(sample) | |
| inputs = {} | |
| for key, value in sample.items(): | |
| if key == "task" or not isinstance(value, torch.Tensor): | |
| continue | |
| if value.dtype == torch.bool: | |
| inputs[key] = value.reshape(1).to("cuda") | |
| elif value.dtype in (torch.int32, torch.int64, torch.int16, torch.int8, torch.uint8): | |
| inputs[key] = value[None].to("cuda") | |
| elif value.is_floating_point(): | |
| inputs[key] = value[None].to(device="cuda", dtype=DTYPE) | |
| else: | |
| inputs[key] = value[None].to("cuda") | |
| return inputs | |
| def _render_action_plot(actions: np.ndarray) -> Image.Image: | |
| """actions: [horizon, 7] -> matplotlib figure image.""" | |
| horizon = actions.shape[0] | |
| steps = np.arange(1, horizon + 1) | |
| fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(11, 4.2)) | |
| # Translation + rotation (absolute end-effector pose). | |
| for i in range(min(6, actions.shape[1])): | |
| ax1.plot(steps, actions[:, i], marker="o", ms=3, label=ACTION_LABELS[i]) | |
| ax1.axhline(0.0, color="0.7", lw=0.8, ls="--") | |
| ax1.set_title("Predicted end-effector pose") | |
| ax1.set_xlabel("action step") | |
| ax1.set_ylabel("value") | |
| ax1.legend(fontsize=8, ncol=2) | |
| ax1.grid(alpha=0.25) | |
| # Gripper command. | |
| if actions.shape[1] >= 7: | |
| ax2.step(steps, actions[:, 6], where="mid", color="#d62728", marker="s", ms=4) | |
| ax2.set_ylim(-1.15, 1.15) | |
| ax2.set_title("Gripper command (+1 open / -1 close)") | |
| ax2.set_xlabel("action step") | |
| ax2.set_ylabel("gripper") | |
| ax2.grid(alpha=0.25) | |
| fig.tight_layout() | |
| tmp = tempfile.NamedTemporaryFile(suffix=".png", delete=False) | |
| fig.savefig(tmp.name, dpi=110) | |
| plt.close(fig) | |
| return Image.open(tmp.name) | |
| def _render_trajectory_3d(actions: np.ndarray) -> Image.Image: | |
| """Plot the predicted absolute 3D end-effector positions as a path.""" | |
| path = actions[:, :3] | |
| fig = plt.figure(figsize=(5.2, 4.6)) | |
| ax = fig.add_subplot(111, projection="3d") | |
| ax.plot(path[:, 0], path[:, 1], path[:, 2], "-o", ms=3, color="#1f77b4") | |
| ax.scatter(path[0, 0], path[0, 1], path[0, 2], color="green", s=50, label="start") | |
| ax.scatter(path[-1, 0], path[-1, 1], path[-1, 2], color="red", s=50, label="end") | |
| ax.set_title("Predicted 3D end-effector trajectory") | |
| ax.set_xlabel("x") | |
| ax.set_ylabel("y") | |
| ax.set_zlabel("z") | |
| ax.legend(fontsize=8) | |
| fig.tight_layout() | |
| tmp = tempfile.NamedTemporaryFile(suffix=".png", delete=False) | |
| fig.savefig(tmp.name, dpi=110) | |
| plt.close(fig) | |
| return Image.open(tmp.name) | |
| def _decode_world_image(recon_images) -> Image.Image | None: | |
| """recon_images: [num_views, 3, H, W] in [-1, 1] -> predicted head-view frame.""" | |
| if recon_images is None: | |
| return None | |
| t = recon_images.detach().float().cpu() | |
| if t.dim() == 4: | |
| t = t[0] # head view | |
| elif t.dim() == 3: | |
| pass | |
| else: | |
| return None | |
| t = (t.clamp(-1, 1) + 1.0) / 2.0 | |
| arr = (t.permute(1, 2, 0).numpy() * 255.0).clip(0, 255).astype(np.uint8) | |
| return Image.fromarray(arr) | |
| def _estimate_duration(head_img=None, wrist_img=None, instruction=None, state_text=None, predict_world=True, *args, **kwargs): | |
| # Measured on ZeroGPU A10G: ~3-4s action-only, ~6-8s with world-model decode. | |
| # First call also pays a one-time lazy CUDA transfer (~3s). Sized tight with | |
| # margin so we don't waste visitor GPU quota. | |
| return 45 if predict_world else 30 | |
| def predict(head_img, wrist_img, instruction, state_text=None, predict_world=True): | |
| """Predict a chunk of robot actions from two camera views and an instruction. | |
| Args: | |
| head_img: RGB image from the head (agentview) camera of the robot scene. | |
| wrist_img: RGB image from the wrist (eye-in-hand) camera. | |
| instruction: natural-language task, e.g. "put the bowl on the plate". | |
| state_text: 8 comma-separated robot proprioceptive state values. | |
| predict_world: also decode the world-model prediction of the next head frame. | |
| Returns: | |
| (action_plot, trajectory_3d, world_image, summary_text) | |
| """ | |
| _ensure_cuda() | |
| if head_img is None or wrist_img is None: | |
| raise gr.Error("Please provide both a head-camera and a wrist-camera image.") | |
| instruction = (instruction or "").strip() | |
| if not instruction: | |
| raise gr.Error("Please enter a task instruction.") | |
| if state_text is None or str(state_text).strip() == "": | |
| state_text = ", ".join(str(x) for x in DEFAULT_STATE) | |
| try: | |
| state = [float(x) for x in str(state_text).replace(";", ",").split(",") if x.strip() != ""] | |
| except ValueError: | |
| raise gr.Error("Robot state must be comma-separated numbers.") | |
| if len(state) < 8: | |
| state = (state + DEFAULT_STATE)[:8] | |
| state = state[:8] | |
| t0 = time.perf_counter() | |
| inputs = _prepare_inputs(head_img, wrist_img, state, instruction) | |
| with torch.no_grad(): | |
| action_pred, recon_images = policy.predict_action_chunk(inputs, decode_image=bool(predict_world)) | |
| # action_pred: [B, chunk, 32] normalized & padded. Slice to the real chunk | |
| # horizon and the 7 LIBERO action dims BEFORE unnormalizing (stats are | |
| # 7-dim); mirrors evaluation/Libero/model_server.py. | |
| action_pred = action_pred[0, :_CHUNK, :TARGET_ACTION_DIM] | |
| action_pred = _unnormalize_action_fn({"action": action_pred})["action"] | |
| actions = action_pred.detach().float().cpu().numpy() | |
| elapsed = time.perf_counter() - t0 | |
| action_plot = _render_action_plot(actions) | |
| traj_3d = _render_trajectory_3d(actions) | |
| world_img = _decode_world_image(recon_images) if predict_world else None | |
| first = actions[0] | |
| summary = ( | |
| f"**Instruction:** {instruction}\n\n" | |
| f"Predicted a chunk of **{actions.shape[0]} actions** " | |
| f"(7-DoF end-effector control) in **{elapsed:.1f}s**.\n\n" | |
| f"**Immediate action (step 1):** " | |
| + ", ".join(f"{ACTION_LABELS[i]}={first[i]:+.3f}" for i in range(min(7, len(first)))) | |
| ) | |
| return action_plot, traj_3d, world_img, summary | |
| # --------------------------------------------------------------------------- | |
| # UI | |
| # --------------------------------------------------------------------------- | |
| CSS = """ | |
| #col-container { max-width: 1180px; margin: 0 auto; } | |
| .dark .gradio-container { color: var(--body-text-color); } | |
| """ | |
| _EX = APP_DIR / "examples" | |
| EXAMPLES = [ | |
| [str(_EX / "bowl_on_plate_head.png"), str(_EX / "bowl_on_plate_wrist.png"), "put the bowl on the plate"], | |
| [str(_EX / "wine_on_rack_head.png"), str(_EX / "wine_on_rack_wrist.png"), "put the wine bottle on the rack"], | |
| [str(_EX / "drawer_bowl_head.png"), str(_EX / "drawer_bowl_wrist.png"), "open the top drawer and put the bowl inside"], | |
| [str(_EX / "cream_cheese_head.png"), str(_EX / "cream_cheese_wrist.png"), "put the cream cheese in the bowl"], | |
| ] | |
| with gr.Blocks(theme=gr.themes.Citrus(), css=CSS) as demo: | |
| with gr.Column(elem_id="col-container"): | |
| gr.Markdown( | |
| """ | |
| # 🤖 WSA-Base — World-Spatial-Action Robot Control | |
| A **3D-centric embodied foundation model** for generalizable robot control | |
| ([paper](https://huggingface.co/papers/2607.03941) · | |
| [code](https://github.com/zaleni/WSA) · | |
| [model](https://huggingface.co/zaleni/WSA-Base-LIBERO)). | |
| Give it a **head-camera** view + **wrist-camera** view of a tabletop robot | |
| scene and a **task instruction**; WSA predicts a short chunk of 7-DoF | |
| end-effector actions and can decode its world-model prediction of the next frame. | |
| Weights: `zaleni/WSA-Base-LIBERO` (3B, Qwen3-VL-2B backbone, LIBERO adapter). | |
| """ | |
| ) | |
| with gr.Row(): | |
| with gr.Column(scale=1): | |
| head_in = gr.Image(label="Head camera (agentview)", type="pil", height=260) | |
| wrist_in = gr.Image(label="Wrist camera (eye-in-hand)", type="pil", height=260) | |
| instruction_in = gr.Textbox( | |
| label="Task instruction", | |
| placeholder="e.g. put the bowl on the plate", | |
| ) | |
| with gr.Accordion("Advanced", open=False): | |
| state_in = gr.Textbox( | |
| label="Robot proprioceptive state (8 comma-separated values)", | |
| value=", ".join(str(x) for x in DEFAULT_STATE), | |
| ) | |
| world_in = gr.Checkbox( | |
| label="Also decode world-model next-frame prediction (slower)", | |
| value=True, | |
| ) | |
| run = gr.Button("Predict actions", variant="primary") | |
| with gr.Column(scale=1): | |
| summary_out = gr.Markdown() | |
| action_out = gr.Image(label="Predicted action chunk", height=300) | |
| traj_out = gr.Image(label="3D trajectory", height=300) | |
| world_out = gr.Image(label="World-model predicted next head frame", height=260) | |
| gr.Examples( | |
| examples=EXAMPLES, | |
| inputs=[head_in, wrist_in, instruction_in], | |
| outputs=[action_out, traj_out, world_out, summary_out], | |
| fn=predict, | |
| cache_examples=False, | |
| run_on_click=True, | |
| ) | |
| run.click( | |
| predict, | |
| inputs=[head_in, wrist_in, instruction_in, state_in, world_in], | |
| outputs=[action_out, traj_out, world_out, summary_out], | |
| api_name="predict", | |
| ) | |
| if __name__ == "__main__": | |
| demo.launch(mcp_server=True) | |