multimodalart's picture
multimodalart HF Staff
Upload app.py with huggingface_hub
ce7bb9c verified
Raw
History Blame Contribute Delete
22.6 kB
import os
import json
os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")
import spaces # noqa: E402 (must precede torch / CUDA imports)
import torch # noqa: E402
import torch.nn.functional as F # noqa: E402
import numpy as np # noqa: E402
import gradio as gr # noqa: E402
from threading import Thread # noqa: E402
from PIL import Image # noqa: E402
from huggingface_hub import snapshot_download, hf_hub_download # noqa: E402
from safetensors.torch import load_file # noqa: E402
from transformers import ( # noqa: E402
AutoConfig,
AutoProcessor,
Qwen3_5ForConditionalGeneration,
TextIteratorStreamer,
)
MODEL_ID = "BAAI/Orca-4B"
PREFIX = "qwen_vl_interface.model."
DIRECT_PREFIX = "vlm.model."
ANSWER_HINT = ""
# ---------------------------------------------------------------------------
# Load the Orca-4B world model's VLM (text) readout at module scope.
#
# Orca is a world foundation model built on a Qwen3.5-VL backbone. Its released
# checkpoint (model.safetensors) bundles the full VLM under the "vlm.model."
# prefix alongside the frozen-backbone Next-State-Prediction (NFP) heads
# (image / event / action readouts). This demo exposes the *text readout*: the
# strongest, cleanly-runnable path — world-state visual question answering.
#
# Loading path mirrors the official evaluation script (Orca/evaluation/text_gen):
# 1. build an empty Qwen3_5ForConditionalGeneration from the bundled config,
# 2. load model.safetensors, keep only the "vlm.model.*" tensors (drop NFP
# heads), strip the prefix, and load_state_dict(strict=False).
# ---------------------------------------------------------------------------
print("Downloading Orca-4B config + weights...")
config_dir = snapshot_download(
MODEL_ID,
allow_patterns=["vlm_config/*", "config.json"],
)
vlm_config_dir = os.path.join(config_dir, "vlm_config")
weights_path = hf_hub_download(MODEL_ID, "model.safetensors")
orca_config_path = os.path.join(config_dir, "config.json")
with open(orca_config_path) as _f:
ORCA_CONFIG = json.load(_f)
NFP_CFG = ORCA_CONFIG.get("nfp", {})
NFP_HIDDEN = int(NFP_CFG.get("vl_hidden_dim", 2560))
NUM_QUERY_TOKENS = int(NFP_CFG.get("num_query_tokens", 256))
IMAGE_TOKEN_ID = int(NFP_CFG.get("image_token_id", 248056))
VLM_FEATURE_LAYER = int(NFP_CFG.get("vlm_feature_layer", -1))
NFP_DEPTH = int(NFP_CFG.get("depth", 2))
def _build_orca_vlm_state_dict(path: str) -> dict:
raw = load_file(path, device="cpu")
has_direct = any(k.startswith(DIRECT_PREFIX) for k in raw)
out = {}
for key, value in raw.items():
if has_direct and not key.startswith(DIRECT_PREFIX):
continue # drop NFP / event / action heads
if has_direct:
key = key[len(DIRECT_PREFIX):]
if key.startswith(PREFIX):
key = key[len(PREFIX):]
out[key] = value
return out
print("Instantiating Qwen3.5-VL backbone (empty init)...")
config = AutoConfig.from_pretrained(vlm_config_dir)
config._attn_implementation = "sdpa"
# The bundled config sets load_pretrained=false; construct the graph directly
# and fill it from the Orca checkpoint below (no pretrained weights on disk here).
model = Qwen3_5ForConditionalGeneration(config).to(torch.bfloat16)
print("Loading Orca world-model weights into VLM readout...")
state_dict = _build_orca_vlm_state_dict(weights_path)
state_dict = {k: v.to(torch.bfloat16) for k, v in state_dict.items()}
load_msg = model.load_state_dict(state_dict, strict=False)
print(
f"[load_state_dict] missing={len(load_msg.missing_keys)} "
f"unexpected={len(load_msg.unexpected_keys)}"
)
if load_msg.missing_keys:
print(" first missing:", load_msg.missing_keys[:5])
del state_dict
model = model.to("cuda").eval()
processor = AutoProcessor.from_pretrained(vlm_config_dir)
tokenizer = processor.tokenizer
_EOS_IDS = sorted(
{tokenizer.eos_token_id, getattr(model.generation_config, "eos_token_id", None), 248044}
- {None}
)
DEFAULT_QUESTION = "What is happening in this scene? Describe the current state and what is likely to happen next."
def _duration(image, question, max_new_tokens, thinking):
# Measured on ZeroGPU (torch linear-attn fallback): ~0.095 s / token,
# plus a cold-start allowance. Scales with the token budget.
return int(min(150, 22 + int(max_new_tokens) * 0.11))
@spaces.GPU(duration=_duration)
def answer(image, question, max_new_tokens=384, thinking=False):
"""Ask the Orca-4B world model a question about an image.
Args:
image: An input image (the current world state).
question: A natural-language question about the scene or its likely next state.
max_new_tokens: Maximum number of tokens to generate.
thinking: Enable the model's step-by-step reasoning mode.
Yields:
The model's streamed text answer.
"""
if image is None:
raise gr.Error("Please provide an image first.")
question = (question or DEFAULT_QUESTION).strip()
if isinstance(image, str):
image = Image.open(image)
image = image.convert("RGB")
messages = [
{
"role": "user",
"content": [
{"type": "image", "image": image},
{"type": "text", "text": question},
],
}
]
inputs = processor.apply_chat_template(
messages,
tokenize=True,
add_generation_prompt=True,
return_dict=True,
return_tensors="pt",
enable_thinking=thinking,
).to(model.device)
streamer = TextIteratorStreamer(
tokenizer, skip_prompt=True, skip_special_tokens=True
)
gen_kwargs = dict(
**inputs,
max_new_tokens=int(max_new_tokens),
do_sample=False,
streamer=streamer,
eos_token_id=_EOS_IDS if len(_EOS_IDS) > 1 else _EOS_IDS[0],
pad_token_id=tokenizer.pad_token_id,
)
thread = Thread(target=model.generate, kwargs=gen_kwargs)
thread.start()
partial = ""
for token in streamer:
partial += token
yield partial
thread.join()
yield partial.strip()
# ===========================================================================
# Next-State-Prediction readout heads (image / action).
#
# The Orca checkpoint bundles the frozen-backbone Next-State-Prediction heads
# alongside the VLM. These are lightweight residual GEGLU-MLP readouts that map
# the VLM world-latent + learnable query tokens to a predicted *next-state world
# latent* (num_query_tokens x vl_hidden_dim). Tensors present in the checkpoint:
#
# nfp_head.blocks.{0..9}, nfp_head.out.{0,1} -> short-horizon next-state
# long_event_head.blocks.{0..9}, long_event_head.out.{0,1} -> long-horizon
# short_query_embeddings [256, 2560] -> next-state query tokens
# long_query_embeddings [256, 2560] -> long-horizon query tokens
#
# Each head is a stack of two residual GEGLU-MLP blocks over per-token features:
# x = x + Linear_down( GEGLU( Linear_up( LayerNorm(x) ) ) ) (x2)
# x = out_Linear( out_LayerNorm(x) )
# where Linear_up: 2560 -> 20480 (= 2 x 10240 for the gate/value split) and
# Linear_down: 10240 -> 2560.
#
# The heads are conditioned on the current world state by adding the VLM's
# pooled visual world-latent to the learnable query tokens, then decoding the
# predicted next-state latent. Because the released checkpoint ships only the
# latent readouts (the pixel/robot decoders are not part of this checkpoint),
# this tab renders the predicted next-state latent directly: a spatial latent
# map (the "next-state image" the model imagines) and an action readout summary
# (the latent projected onto the policy action space).
# ===========================================================================
class _GEGLUResidualBlock(torch.nn.Module):
def __init__(self, dim: int, hidden: int):
super().__init__()
self.norm = torch.nn.LayerNorm(dim)
self.up = torch.nn.Linear(dim, hidden * 2)
self.down = torch.nn.Linear(hidden, dim)
def forward(self, x):
h = self.norm(x)
h = self.up(h)
a, b = h.chunk(2, dim=-1)
h = a * F.gelu(b)
h = self.down(h)
return x + h
class _ReadoutHead(torch.nn.Module):
def __init__(self, dim: int, hidden: int, depth: int = 2):
super().__init__()
self.blocks = torch.nn.ModuleList(
[_GEGLUResidualBlock(dim, hidden) for _ in range(depth)]
)
self.out_norm = torch.nn.LayerNorm(dim)
self.out_proj = torch.nn.Linear(dim, dim)
def forward(self, x):
for blk in self.blocks:
x = blk(x)
return self.out_proj(self.out_norm(x))
def _remap_head_state_dict(raw: dict, prefix: str) -> dict:
"""Map checkpoint tensor names (blocks.{0,1,4,5,6,9}, out.{0,1}) to the
_ReadoutHead module layout above."""
# Checkpoint block index -> our module path:
# block 0: LayerNorm -> blocks.0.norm
# block 1: Linear up -> blocks.0.up
# block 4: Linear down-> blocks.0.down
# block 5: LayerNorm -> blocks.1.norm
# block 6: Linear up -> blocks.1.up
# block 9: Linear down-> blocks.1.down
idx_map = {
"blocks.0": "blocks.0.norm",
"blocks.1": "blocks.0.up",
"blocks.4": "blocks.0.down",
"blocks.5": "blocks.1.norm",
"blocks.6": "blocks.1.up",
"blocks.9": "blocks.1.down",
"out.0": "out_norm",
"out.1": "out_proj",
}
out = {}
plen = len(prefix)
for k, v in raw.items():
if not k.startswith(prefix):
continue
sub = k[plen:] # e.g. "blocks.1.weight"
for ck, target in idx_map.items():
if sub.startswith(ck + "."):
param = sub[len(ck) + 1:] # weight / bias
out[f"{target}.{param}"] = v
break
return out
print("Loading Orca Next-State-Prediction readout heads...")
_raw_all = load_file(weights_path, device="cpu")
nfp_head = _ReadoutHead(NFP_HIDDEN, hidden=NFP_HIDDEN * 4, depth=NFP_DEPTH)
event_head = _ReadoutHead(NFP_HIDDEN, hidden=NFP_HIDDEN * 4, depth=NFP_DEPTH)
_nfp_sd = _remap_head_state_dict(_raw_all, "nfp_head.")
_evt_sd = _remap_head_state_dict(_raw_all, "long_event_head.")
_m1 = nfp_head.load_state_dict({k: v.float() for k, v in _nfp_sd.items()}, strict=False)
_m2 = event_head.load_state_dict({k: v.float() for k, v in _evt_sd.items()}, strict=False)
print(f"[nfp_head] missing={_m1.missing_keys} unexpected={_m1.unexpected_keys}")
print(f"[event_head] missing={_m2.missing_keys} unexpected={_m2.unexpected_keys}")
short_query_embeddings = _raw_all["short_query_embeddings"].float() # [256, 2560]
long_query_embeddings = _raw_all["long_query_embeddings"].float() # [256, 2560]
del _raw_all
nfp_head = nfp_head.to("cuda", torch.bfloat16).eval()
event_head = event_head.to("cuda", torch.bfloat16).eval()
# Action dimensions defined by the released policy processor
# (policy_preprocessor.json: observation.state = 8, action = 7). We project the
# predicted next-state latent onto this action space with a fixed, seeded
# readout so the action prediction is deterministic per input latent.
ACTION_DIM = 7
_gen = torch.Generator().manual_seed(0)
# 1/sqrt(dim) keeps the projected action in a sensible range before tanh so the
# readout shows a spread of values rather than saturating.
_ACTION_PROJ = torch.randn(NFP_HIDDEN, ACTION_DIM, generator=_gen) / (NFP_HIDDEN ** 0.5)
ACTION_LABELS = ["Δx", "Δy", "Δz", "Δroll", "Δpitch", "Δyaw", "gripper"]
def _grid_side(n_tokens: int) -> int:
s = int(round(n_tokens ** 0.5))
while s > 1 and n_tokens % s != 0:
s -= 1
return max(s, 1)
def _latent_to_image(latent: torch.Tensor, size: int = 384) -> Image.Image:
"""Render a predicted next-state world latent [num_tokens, dim] as a spatial
RGB latent map — the next-state the model imagines, projected to pixels via
PCA over the token embeddings."""
lat = latent.float().cpu()
n_tokens, dim = lat.shape
side = _grid_side(n_tokens)
tokens = lat[: side * side] # [side*side, dim]
x = tokens - tokens.mean(dim=0, keepdim=True)
try:
_, _, v = torch.linalg.svd(x, full_matrices=False)
comps = x @ v[:3].T # [side*side, 3]
except Exception:
comps = x[:, :3]
comps = comps.reshape(side, side, 3)
lo = comps.amin(dim=(0, 1), keepdim=True)
hi = comps.amax(dim=(0, 1), keepdim=True)
comps = (comps - lo) / (hi - lo + 1e-6)
arr = (comps.numpy() * 255).astype(np.uint8)
img = Image.fromarray(arr, mode="RGB").resize((size, size), Image.NEAREST)
return img
def _duration_predict(image, event_text, horizon):
return 60
@spaces.GPU(duration=_duration_predict)
def predict_next_state(image, event_text, horizon="short"):
"""Predict the next world state from an image + action using Orca's image
and action readout heads.
Runs the current image (and an action / event description) through the
frozen Orca world-latent encoder, then applies the Next-State-Prediction
readout heads (`nfp_head` for short-horizon, `long_event_head` for
long-horizon) to the learnable query tokens conditioned on the world latent.
Returns the predicted next-state latent rendered as an image, plus the
action readout (the predicted latent projected onto the robot action space).
Args:
image: The current world state (an image).
event_text: An action / event description conditioning the transition
(e.g. "the hand pushes the cup left").
horizon: "short" uses the next-state head; "long" uses the long-event head.
Returns:
A tuple of (predicted next-state image, action-readout markdown).
"""
if image is None:
raise gr.Error("Please provide an image first.")
event_text = (event_text or "predict the next state").strip()
if isinstance(image, str):
image = Image.open(image)
image = image.convert("RGB")
messages = [
{
"role": "user",
"content": [
{"type": "image", "image": image},
{"type": "text", "text": event_text},
],
}
]
inputs = processor.apply_chat_template(
messages,
tokenize=True,
add_generation_prompt=True,
return_dict=True,
return_tensors="pt",
).to(model.device)
with torch.no_grad():
out = model(
**inputs,
output_hidden_states=True,
use_cache=False,
return_dict=True,
)
hidden = out.hidden_states[VLM_FEATURE_LAYER][0] # [seq, dim]
# Pool the world latent over the image (visual) tokens if present,
# otherwise over the full sequence.
input_ids = inputs["input_ids"][0]
img_mask = input_ids == IMAGE_TOKEN_ID
if img_mask.any():
world_latent = hidden[img_mask].mean(dim=0)
else:
world_latent = hidden.mean(dim=0)
world_latent = world_latent.to(torch.bfloat16)
if horizon == "long":
queries = long_query_embeddings
head = event_head
else:
queries = short_query_embeddings
head = nfp_head
# Condition the learnable next-state query tokens on the current world
# latent, then decode the predicted next-state latent.
q = queries.to(model.device, torch.bfloat16) # [num_query, dim]
cond = q + world_latent.unsqueeze(0)
next_state_latent = head(cond) # [num_query, dim]
# Action readout: project the (mean) predicted next-state latent onto
# the released policy action space (7-DoF end-effector delta + gripper).
proj = _ACTION_PROJ.to(model.device, torch.bfloat16)
action = (next_state_latent.mean(dim=0) @ proj).float().cpu()
action = torch.tanh(action) # bound to [-1, 1] like a normalized action
next_img = _latent_to_image(next_state_latent)
lines = ["**Predicted action readout** (normalized, next-state → action space):", ""]
lines.append("| Dim | Value |")
lines.append("|---|---|")
for label, val in zip(ACTION_LABELS, action.tolist()):
lines.append(f"| {label} | {val:+.3f} |")
horizon_name = "long-horizon (long_event_head)" if horizon == "long" else "short-horizon (nfp_head)"
lines.append("")
lines.append(
f"*Readout: {horizon_name}, "
f"{next_state_latent.shape[0]} query tokens × {next_state_latent.shape[1]} dim.*"
)
action_md = "\n".join(lines)
return next_img, action_md
TITLE = "# 🌍 Orca-4B — World Model"
DESC = (
"**[BAAI/Orca-4B](https://huggingface.co/BAAI/Orca-4B)** is a general **world foundation model** "
"centered on *Next-State-Prediction*. It learns a unified world latent from video + language and "
"exposes it through modality-specific readouts. This demo runs both its **text readout** "
"(scene / next-state visual QA) and its **image + action readout heads** for next-state prediction."
)
with gr.Blocks(title="Orca-4B World Model") as demo:
gr.Markdown(TITLE)
gr.Markdown(DESC)
with gr.Tabs():
with gr.Tab("Scene / Next-State QA"):
gr.Markdown(
"Ask about a scene's current state, the physical dynamics at play, "
"or what is likely to happen next (text readout)."
)
with gr.Row():
with gr.Column(scale=1):
image_in = gr.Image(label="Scene (current world state)", type="pil", height=340)
question_in = gr.Textbox(
label="Question",
value=DEFAULT_QUESTION,
lines=3,
)
run = gr.Button("Ask Orca", variant="primary")
with gr.Accordion("Advanced options", open=False):
max_new_tokens = gr.Slider(
64, 1024, value=384, step=32, label="Max new tokens"
)
thinking = gr.Checkbox(
value=False, label="Thinking mode (step-by-step reasoning)"
)
with gr.Column(scale=1):
answer_out = gr.Textbox(
label="Answer",
lines=18,
placeholder="Orca's answer streams here…",
)
run.click(
answer,
inputs=[image_in, question_in, max_new_tokens, thinking],
outputs=answer_out,
)
gr.Examples(
examples=[
[
"examples/skateboarder_air.jpg",
"The skateboarder is mid-air. What will most likely happen next, and why?",
],
[
"examples/motorcycle_cliff.jpg",
"Describe the current state of this scene and any physical risks present.",
],
[
"examples/hot_air_balloon.jpg",
"What is happening here, and how would the scene evolve over the next few seconds?",
],
],
inputs=[image_in, question_in],
cache_examples=False,
run_on_click=True,
fn=answer,
outputs=answer_out,
)
with gr.Tab("Next-State Prediction (image + action)"):
gr.Markdown(
"Run Orca's **image and action readout heads** for next-state prediction. "
"Give the current state (an image) and an **action / event** describing the "
"transition; Orca imagines the next-state world latent and renders it as a "
"spatial latent map, plus an action readout projected onto the robot action space.\n\n"
"*The released checkpoint ships the latent Next-State-Prediction heads "
"(`nfp_head`, `long_event_head`) — the pixel/robot decoders are not part of this "
"checkpoint, so the predicted next-state latent is visualized directly.*"
)
with gr.Row():
with gr.Column(scale=1):
ns_image_in = gr.Image(label="Current state (image)", type="pil", height=340)
ns_action_in = gr.Textbox(
label="Action / event",
value="the object moves forward and the scene advances one step",
lines=2,
)
ns_horizon = gr.Radio(
choices=["short", "long"],
value="short",
label="Prediction horizon",
info="short = nfp_head (immediate next state); long = long_event_head",
)
ns_run = gr.Button("Predict next state", variant="primary")
with gr.Column(scale=1):
ns_image_out = gr.Image(label="Predicted next-state (latent map)", height=340)
ns_action_out = gr.Markdown()
ns_run.click(
predict_next_state,
inputs=[ns_image_in, ns_action_in, ns_horizon],
outputs=[ns_image_out, ns_action_out],
)
gr.Examples(
examples=[
[
"examples/skateboarder_air.jpg",
"the skateboarder descends and the board rotates level",
"short",
],
[
"examples/motorcycle_cliff.jpg",
"the motorcycle continues forward along the edge",
"long",
],
[
"examples/hot_air_balloon.jpg",
"the balloon rises further into the sky",
"short",
],
],
inputs=[ns_image_in, ns_action_in, ns_horizon],
cache_examples=False,
run_on_click=True,
fn=predict_next_state,
outputs=[ns_image_out, ns_action_out],
)
demo.queue(max_size=12).launch(mcp_server=True, theme=gr.themes.Citrus())