multimodalart's picture
multimodalart HF Staff
Upload app.py with huggingface_hub
71ff53d verified
Raw History Blame Contribute Delete
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
@spaces.GPU(duration=_estimate_duration)
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)