Spaces:
Running on Zero
Running on Zero
| """DeepThinkVLA demo — chain-of-thought reasoning + robot action chunks. | |
| Faithful port of the authors' single-step inference path | |
| (`src/experiments/deepthinkvla_utils.py::get_vla_action` in | |
| https://github.com/OpenBMB/DeepThinkVLA) to a Gradio Space. | |
| """ | |
| import os | |
| os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True") | |
| import spaces # noqa: E402 (must precede torch / CUDA-touching imports) | |
| import io # noqa: E402 | |
| import json # noqa: E402 | |
| import time # noqa: E402 | |
| import gradio as gr # noqa: E402 | |
| import matplotlib # noqa: E402 | |
| matplotlib.use("Agg") | |
| import matplotlib.pyplot as plt # noqa: E402 | |
| import numpy as np # noqa: E402 | |
| import torch # noqa: E402 | |
| from huggingface_hub import snapshot_download # noqa: E402 | |
| from PIL import Image # noqa: E402 | |
| from transformers import AutoProcessor, GenerationConfig # noqa: E402 | |
| from dt_datasets.normalize import Unnormalize_Action # noqa: E402 | |
| from sft.constants import ( # noqa: E402 | |
| ACTION_DIM, | |
| ACTION_MASK, | |
| ACTION_PROPRIO_NORMALIZATION_TYPE, | |
| NUM_ACTIONS_CHUNK, | |
| ) | |
| from sft.modeling_deepthinkvla import DeepThinkVLA # noqa: E402 | |
| # ---------------------------------------------------------------------------- | |
| # Constants (copied verbatim from the reference eval code) | |
| # ---------------------------------------------------------------------------- | |
| MODEL_ID = "yinchenghust/deepthinkvla_libero_cot_rl" | |
| THINK_PREFIX = ( | |
| "First output the thinking process in <think></think> tags and then output " | |
| "the final action in <action></action>." | |
| ) | |
| DEEPTHINKVLA_IMAGE_SIZE = 224 | |
| DIM_LABELS = ["dx", "dy", "dz", "d_roll", "d_pitch", "d_yaw", "gripper"] | |
| # ---------------------------------------------------------------------------- | |
| # Load model / processor / action de-normalizer | |
| # ---------------------------------------------------------------------------- | |
| print(f"Downloading {MODEL_ID} …", flush=True) | |
| CKPT_DIR = snapshot_download(MODEL_ID) | |
| processor = AutoProcessor.from_pretrained(CKPT_DIR) | |
| model = DeepThinkVLA.from_pretrained( | |
| CKPT_DIR, | |
| torch_dtype=torch.bfloat16, | |
| attn_implementation="sdpa", | |
| ) | |
| model.eval() | |
| model = model.to("cuda") | |
| with open(os.path.join(CKPT_DIR, "norm_stats.json")) as f: | |
| _norm_stats = json.load(f) | |
| for _k in _norm_stats["action"]: | |
| _norm_stats["action"][_k] = np.array(_norm_stats["action"][_k], dtype=np.float64) | |
| unnormalize_action = Unnormalize_Action( | |
| normalization_type=ACTION_PROPRIO_NORMALIZATION_TYPE, | |
| stats=_norm_stats["action"], | |
| action_mask=ACTION_MASK, | |
| ) | |
| print("Model ready.", flush=True) | |
| # ---------------------------------------------------------------------------- | |
| # Pre / post processing | |
| # ---------------------------------------------------------------------------- | |
| def _prepare_image(img) -> Image.Image: | |
| """np.uint8 (H, W, 3) -> 224x224 RGB PIL (bilinear, as in the reference).""" | |
| if img is None: | |
| raise gr.Error("Both a third-person image and a wrist image are required.") | |
| pil = Image.fromarray(np.asarray(img, dtype=np.uint8)).convert("RGB") | |
| if pil.size != (DEEPTHINKVLA_IMAGE_SIZE, DEEPTHINKVLA_IMAGE_SIZE): | |
| pil = pil.resize( | |
| (DEEPTHINKVLA_IMAGE_SIZE, DEEPTHINKVLA_IMAGE_SIZE), Image.BILINEAR | |
| ) | |
| return pil | |
| def _binarize_gripper(actions: np.ndarray) -> np.ndarray: | |
| out = actions.copy() | |
| out[..., -1] = np.sign(out[..., -1]) | |
| return out | |
| def render_action_plot(actions: np.ndarray) -> Image.Image: | |
| """Plot the action chunk: cumulative EE path + per-DoF deltas.""" | |
| a = np.asarray(actions, dtype=np.float64) | |
| n = a.shape[0] | |
| steps = np.arange(1, n + 1) | |
| path = np.vstack([np.zeros((1, 3)), np.cumsum(a[:, :3], axis=0)]) | |
| grip = np.sign(a[:, 6]) | |
| fig = plt.figure(figsize=(15.0, 4.4), dpi=110) | |
| # --- 3D cumulative end-effector displacement --------------------------- | |
| ax = fig.add_subplot(1, 3, 1, projection="3d") | |
| ax.plot(path[:, 0], path[:, 1], path[:, 2], color="#4b5563", lw=1.4, zorder=1) | |
| sc = ax.scatter( | |
| path[1:, 0], path[1:, 1], path[1:, 2], c=steps, cmap="viridis", s=46, zorder=2 | |
| ) | |
| ax.scatter(0, 0, 0, marker="o", s=70, facecolors="none", edgecolors="k", lw=1.4) | |
| closed = grip > 0 | |
| if closed.any(): | |
| ax.scatter( | |
| path[1:, 0][closed], | |
| path[1:, 1][closed], | |
| path[1:, 2][closed], | |
| marker="x", | |
| s=90, | |
| c="crimson", | |
| label="gripper closing", | |
| ) | |
| ax.legend(loc="upper left", fontsize=8) | |
| ax.set_title("Cumulative EE displacement\n(open circle = current pose)", fontsize=10) | |
| ax.set_xlabel("x", fontsize=9) | |
| ax.set_ylabel("y", fontsize=9) | |
| ax.set_zlabel("z", fontsize=9) | |
| ax.tick_params(labelsize=7) | |
| cb = fig.colorbar(sc, ax=ax, pad=0.12, shrink=0.7) | |
| cb.set_label("step", fontsize=8) | |
| cb.ax.tick_params(labelsize=7) | |
| # --- translation deltas ------------------------------------------------ | |
| ax2 = fig.add_subplot(1, 3, 2) | |
| for i, (lbl, color) in enumerate(zip(DIM_LABELS[:3], ["#2563eb", "#16a34a", "#db2777"])): | |
| ax2.plot(steps, a[:, i], marker="o", ms=4, lw=1.6, color=color, label=lbl) | |
| ax2.axhline(0.0, color="#9ca3af", lw=0.8, ls="--") | |
| ax2.set_title("Translation deltas per step", fontsize=10) | |
| ax2.set_xlabel("step in chunk", fontsize=9) | |
| ax2.set_ylabel("delta position (OSC_POSE units)", fontsize=9) | |
| ax2.set_xticks(steps) | |
| ax2.tick_params(labelsize=8) | |
| ax2.legend(fontsize=8) | |
| ax2.grid(alpha=0.25) | |
| # --- rotation deltas + gripper ---------------------------------------- | |
| ax3 = fig.add_subplot(1, 3, 3) | |
| for i, (lbl, color) in enumerate( | |
| zip(DIM_LABELS[3:6], ["#7c3aed", "#f59e0b", "#0891b2"]), start=3 | |
| ): | |
| ax3.plot(steps, a[:, i], marker="o", ms=4, lw=1.6, color=color, label=lbl) | |
| ax3.axhline(0.0, color="#9ca3af", lw=0.8, ls="--") | |
| ax3.set_title("Rotation deltas + gripper command", fontsize=10) | |
| ax3.set_xlabel("step in chunk", fontsize=9) | |
| ax3.set_ylabel("delta rotation (axis-angle)", fontsize=9) | |
| ax3.set_xticks(steps) | |
| ax3.tick_params(labelsize=8) | |
| ax3.grid(alpha=0.25) | |
| ax4 = ax3.twinx() | |
| ax4.step(steps, grip, where="mid", color="crimson", lw=1.8, label="gripper (+1 close)") | |
| ax4.set_ylim(-1.6, 1.6) | |
| ax4.set_yticks([-1, 1]) | |
| ax4.set_ylabel("gripper", fontsize=9, color="crimson") | |
| ax4.tick_params(labelsize=8, colors="crimson") | |
| h1, l1 = ax3.get_legend_handles_labels() | |
| h2, l2 = ax4.get_legend_handles_labels() | |
| ax3.legend(h1 + h2, l1 + l2, fontsize=8, loc="upper right") | |
| fig.tight_layout() | |
| buf = io.BytesIO() | |
| fig.savefig(buf, format="png", bbox_inches="tight") | |
| plt.close(fig) | |
| buf.seek(0) | |
| return Image.open(buf).convert("RGB") | |
| # ---------------------------------------------------------------------------- | |
| # Inference | |
| # ---------------------------------------------------------------------------- | |
| def predict( | |
| third_person_image, | |
| wrist_image, | |
| instruction: str, | |
| max_cot_tokens: int = 512, | |
| ): | |
| """Predict a chain-of-thought trace and a 10-step robot action chunk. | |
| Args: | |
| third_person_image: agent-view RGB observation of the tabletop scene. | |
| wrist_image: eye-in-hand RGB observation from the gripper camera. | |
| instruction: natural-language task, e.g. "pick up the alphabet soup and place it in the basket". | |
| max_cot_tokens: cap on the number of chain-of-thought tokens to generate. | |
| """ | |
| if not instruction or not instruction.strip(): | |
| raise gr.Error("Please provide a task instruction.") | |
| images = [_prepare_image(third_person_image), _prepare_image(wrist_image)] | |
| image_token = processor.tokenizer.additional_special_tokens[0] | |
| prompt = ( | |
| image_token * len(images) | |
| + THINK_PREFIX | |
| + f"Task: {instruction.strip().lower()};" | |
| ) | |
| inputs = processor(text=[prompt], images=images, return_tensors="pt").to( | |
| "cuda", dtype=torch.bfloat16 | |
| ) | |
| generation_config = GenerationConfig( | |
| max_new_tokens=int(max_cot_tokens), | |
| do_sample=False, | |
| pad_token_id=processor.tokenizer.pad_token_id, | |
| bos_token_id=processor.tokenizer.bos_token_id, | |
| eos_token_id=None, | |
| use_cache=True, | |
| num_beams=1, | |
| temperature=None, | |
| top_p=None, | |
| top_k=None, | |
| ) | |
| t0 = time.time() | |
| with torch.inference_mode(): | |
| normalized_actions, input_cot_ids = model.predict_cot_action( | |
| input_ids=inputs["input_ids"], | |
| pixel_values=inputs["pixel_values"], | |
| attention_mask=inputs["attention_mask"], | |
| generation_config=generation_config, | |
| ) | |
| elapsed = time.time() - t0 | |
| assert normalized_actions.shape == (NUM_ACTIONS_CHUNK, ACTION_DIM) | |
| actions = unnormalize_action(torch.from_numpy(normalized_actions)).numpy() | |
| actions = _binarize_gripper(actions) | |
| n_new = int(input_cot_ids.shape[-1] - inputs["input_ids"].shape[-1]) | |
| cot_text = processor.tokenizer.decode( | |
| input_cot_ids[0, inputs["input_ids"].shape[-1] : -1] | |
| ) | |
| print( | |
| f"[predict] cot_tokens={n_new} chunk={actions.shape} " | |
| f"latency={elapsed:.2f}s", | |
| flush=True, | |
| ) | |
| if "</think>" not in cot_text: | |
| cot_text += ( | |
| "\n\n[warning] the reasoning trace hit the token cap before closing " | |
| "</think>; raise 'Max CoT tokens' for a complete trace." | |
| ) | |
| table = [ | |
| [i + 1] + [round(float(v), 4) for v in actions[i]] for i in range(actions.shape[0]) | |
| ] | |
| plot = render_action_plot(actions) | |
| summary = ( | |
| f"**{actions.shape[0]} x {actions.shape[1]} action chunk** — " | |
| f"{n_new} reasoning tokens generated in {elapsed:.1f}s. " | |
| f"Net displacement (x, y, z) = " | |
| f"({actions[:, 0].sum():+.3f}, {actions[:, 1].sum():+.3f}, {actions[:, 2].sum():+.3f}); " | |
| f"gripper ends {'closed' if actions[-1, 6] > 0 else 'open'}." | |
| ) | |
| return cot_text, plot, table, summary | |
| # ---------------------------------------------------------------------------- | |
| # UI | |
| # ---------------------------------------------------------------------------- | |
| CSS = """ | |
| #col-container { max-width: 1200px; margin: 0 auto; } | |
| .dark .gradio-container { color: var(--body-text-color); } | |
| """ | |
| HEADER = """# DeepThinkVLA — reason, then act | |
| <a href="https://huggingface.co/papers/2511.15669">Paper</a> · | |
| <a href="https://github.com/OpenBMB/DeepThinkVLA">Code</a> · | |
| <a href="https://huggingface.co/yinchenghust/deepthinkvla_libero_cot_rl">Model</a> | |
| A 3B PaliGemma-based Vision-Language-Action model trained with SFT + RL so that its | |
| chain-of-thought actually *helps* the action it emits. Give it a tabletop scene | |
| (agent view + wrist camera) and a task; it writes out its reasoning, then predicts the | |
| next **10-step, 7-DoF action chunk** in one non-autoregressive pass. | |
| """ | |
| NOTES = """ | |
| **Reading the output.** Actions are LIBERO `OSC_POSE` commands: three normalized | |
| end-effector position deltas, three axis-angle rotation deltas, and a binary gripper | |
| command (`+1` closing, `-1` opening). At full scale one step is roughly 5 cm / 0.5 rad. | |
| **About the images.** DeepThinkVLA is trained on LIBERO renders that are rotated 180° | |
| by the standard OpenVLA data pipeline, so the example frames look mirrored — that is | |
| exactly what the policy expects. Feeding it ordinary photographs is out of distribution. | |
| Example frames come from the authors' [`yinchenghust/libero_cot`](https://huggingface.co/datasets/yinchenghust/libero_cot) | |
| dataset (Apache-2.0). Model code vendored from OpenBMB/DeepThinkVLA (MIT). | |
| """ | |
| with gr.Blocks(title="DeepThinkVLA") as demo: | |
| with gr.Column(elem_id="col-container"): | |
| gr.Markdown(HEADER) | |
| with gr.Row(): | |
| with gr.Column(scale=1): | |
| third = gr.Image(label="Agent-view image", type="numpy", height=240) | |
| wrist = gr.Image(label="Wrist-camera image", type="numpy", height=240) | |
| with gr.Column(scale=2): | |
| instruction = gr.Textbox( | |
| label="Task instruction", | |
| placeholder="pick up the alphabet soup and place it in the basket", | |
| lines=2, | |
| ) | |
| run = gr.Button("Reason and predict actions", variant="primary") | |
| cot = gr.Textbox( | |
| label="Chain-of-thought", | |
| lines=11, | |
| interactive=False, | |
| ) | |
| summary = gr.Markdown() | |
| plot = gr.Image(label="Predicted action chunk", type="pil", height=330) | |
| table = gr.Dataframe( | |
| label="Action chunk (10 steps x 7 DoF)", | |
| headers=["step"] + DIM_LABELS, | |
| datatype=["number"] * 8, | |
| interactive=False, | |
| ) | |
| with gr.Accordion("Advanced settings", open=False): | |
| max_cot = gr.Slider( | |
| label="Max CoT tokens", | |
| minimum=64, | |
| maximum=1024, | |
| step=32, | |
| value=512, | |
| ) | |
| gr.Examples( | |
| examples=[ | |
| [ | |
| "examples/alphabet_soup_third.png", | |
| "examples/alphabet_soup_wrist.png", | |
| "pick up the alphabet soup and place it in the basket", | |
| ], | |
| [ | |
| "examples/middle_drawer_third.png", | |
| "examples/middle_drawer_wrist.png", | |
| "open the middle drawer of the cabinet", | |
| ], | |
| [ | |
| "examples/black_bowl_third.png", | |
| "examples/black_bowl_wrist.png", | |
| "pick up the black bowl between the plate and the ramekin and place it on the plate", | |
| ], | |
| [ | |
| "examples/moka_pots_third.png", | |
| "examples/moka_pots_wrist.png", | |
| "put both moka pots on the stove", | |
| ], | |
| ], | |
| inputs=[third, wrist, instruction], | |
| outputs=[cot, plot, table, summary], | |
| fn=predict, | |
| cache_examples=True, | |
| cache_mode="lazy", | |
| ) | |
| gr.Markdown(NOTES) | |
| run.click( | |
| predict, | |
| inputs=[third, wrist, instruction, max_cot], | |
| outputs=[cot, plot, table, summary], | |
| api_name="predict", | |
| ) | |
| if __name__ == "__main__": | |
| demo.queue().launch(theme=gr.themes.Citrus(), css=CSS) | |