multimodalart's picture
multimodalart HF Staff
Upload folder using huggingface_hub
7c21e4d verified
Raw
History Blame Contribute Delete
22 kB
import os
os.environ["PYTORCH_CUDA_ALLOC_CONF"] = "expandable_segments:True"
import spaces
import sys
import time
import tempfile
import numpy as np
import torch
import torch.nn.functional as F
import gradio as gr
from PIL import Image
from einops import rearrange
from tqdm import tqdm
# ------------------------------------------------------------------ #
# Model + config setup (module scope, eagerly on GPU)
# ------------------------------------------------------------------ #
MODEL_ID = "gangweix/next-forcing-base"
from huggingface_hub import snapshot_download
_model_path = snapshot_download(
MODEL_ID,
repo_type="model",
allow_patterns=[
"transformer/*",
"vae/*",
"text_encoder/*",
"tokenizer/*",
],
)
from wan_va.modules.utils import (
WanVAEStreamingWrapper,
load_text_encoder,
load_tokenizer,
load_transformer,
load_vae,
)
from wan_va.utils.scheduler import FlowMatchScheduler
from wan_va.utils.utils import get_mesh_id, data_seq_to_patch
DTYPE = torch.bfloat16
DEVICE = "cuda"
# ---- Demo config (matches va_demo_cfg.py) ----
CONFIG = dict(
attn_window=30,
frame_chunk_size=4,
env_type="none",
height=256,
width=256,
action_dim=30,
action_per_frame=8,
obs_cam_keys=["observation.images.top", "observation.images.wrist"],
guidance_scale=5,
action_guidance_scale=1,
num_inference_steps=5,
video_exec_step=-1,
action_num_inference_steps=10,
snr_shift=5.0,
action_snr_shift=1.0,
patch_size=(1, 2, 2),
used_action_channel_ids=list(range(0, 5)) + list(range(28, 29)),
action_norm_method="quantiles",
norm_stat={
"q01": [
-90.60303497314453,
-98.73043060302734,
-79.9008560180664,
48.95470428466797,
-32.794578552246094,
] + [0.0] * 23 + [0.8250824809074402, 0],
"q99": [
71.735107421875,
65.89081573486328,
92.87967681884766,
100.0,
22.784151077270508,
] + [0.0] * 23 + [100.0, 0],
},
)
# Inverse action channel mapping
inverse_used_action_channel_ids = [len(CONFIG["used_action_channel_ids"])] * CONFIG[
"action_dim"
]
for i, j in enumerate(CONFIG["used_action_channel_ids"]):
inverse_used_action_channel_ids[j] = i
CONFIG["inverse_used_action_channel_ids"] = inverse_used_action_channel_ids
# ---- Load model components ----
vae = load_vae(os.path.join(_model_path, "vae"), torch_dtype=DTYPE, torch_device=DEVICE)
streaming_vae = WanVAEStreamingWrapper(vae)
tokenizer = load_tokenizer(os.path.join(_model_path, "tokenizer"))
text_encoder = load_text_encoder(
os.path.join(_model_path, "text_encoder"), torch_dtype=DTYPE, torch_device=DEVICE
)
transformer = load_transformer(
os.path.join(_model_path, "transformer"),
torch_dtype=DTYPE,
torch_device=DEVICE,
attn_mode="torch",
disable_mcp=True,
)
transformer.eval().requires_grad_(False)
scheduler = FlowMatchScheduler(shift=CONFIG["snr_shift"], sigma_min=0.0, extra_one_step=True)
action_scheduler = FlowMatchScheduler(
shift=CONFIG["action_snr_shift"], sigma_min=0.0, extra_one_step=True
)
scheduler.set_timesteps(1000, training=True)
action_scheduler.set_timesteps(1000, training=True)
action_mask = torch.zeros([CONFIG["action_dim"]]).bool()
action_mask[CONFIG["used_action_channel_ids"]] = True
actions_q01 = torch.tensor(CONFIG["norm_stat"]["q01"], dtype=torch.float32).reshape(-1, 1, 1)
actions_q99 = torch.tensor(CONFIG["norm_stat"]["q99"], dtype=torch.float32).reshape(-1, 1, 1)
from diffusers.video_processor import VideoProcessor
video_processor = VideoProcessor(vae_scale_factor=1)
# ------------------------------------------------------------------ #
# Inference helpers
# ------------------------------------------------------------------ #
def _get_t5_prompt_embeds(prompt, max_sequence_length=512):
from diffusers.pipelines.wan.pipeline_wan import prompt_clean
prompt_list = [prompt] if isinstance(prompt, str) else prompt
prompt_list = [prompt_clean(u) for u in prompt_list]
batch_size = len(prompt_list)
text_inputs = tokenizer(
prompt_list,
padding="max_length",
max_length=max_sequence_length,
truncation=True,
add_special_tokens=True,
return_attention_mask=True,
return_tensors="pt",
)
text_input_ids, mask = text_inputs.input_ids, text_inputs.attention_mask
seq_lens = mask.gt(0).sum(dim=1).long()
text_encoder_device = next(text_encoder.parameters()).device
prompt_embeds = text_encoder(
text_input_ids.to(text_encoder_device), mask.to(text_encoder_device)
).last_hidden_state
prompt_embeds = prompt_embeds.to(dtype=DTYPE, device=DEVICE)
prompt_embeds = [u[:v] for u, v in zip(prompt_embeds, seq_lens)]
prompt_embeds = torch.stack(
[
torch.cat([u, u.new_zeros(max_sequence_length - u.size(0), u.size(1))])
for u in prompt_embeds
],
dim=0,
)
_, seq_len, _ = prompt_embeds.shape
prompt_embeds = prompt_embeds.repeat(1, 1, 1)
prompt_embeds = prompt_embeds.view(batch_size, seq_len, -1)
return prompt_embeds.to(DEVICE)
def encode_prompt(prompt):
prompt_embeds = _get_t5_prompt_embeds(prompt)
neg_prompt_embeds = _get_t5_prompt_embeds("")
return prompt_embeds, neg_prompt_embeds
def normalize_latents(latents, latents_mean, latents_std):
latents_mean = latents_mean.view(1, -1, 1, 1, 1).to(device=latents.device)
latents_std = latents_std.view(1, -1, 1, 1, 1).to(device=latents.device)
latents = ((latents.float() - latents_mean) * latents_std).to(latents)
return latents
def encode_obs(images_dict_list):
"""Encode observation images into latent space.
Args:
images_dict_list: list of dicts, each mapping cam_key -> np.ndarray(H,W,3) uint8
"""
images = images_dict_list
if not isinstance(images, list):
images = [images]
if len(images) < 1:
return None
videos = []
for k_i, k in enumerate(CONFIG["obs_cam_keys"]):
height_i, width_i = CONFIG["height"], CONFIG["width"]
history_video_k = (
torch.from_numpy(np.stack([each[k] for each in images]))
.float()
.permute(3, 0, 1, 2)
)
history_video_k = F.interpolate(
history_video_k,
size=(height_i, width_i),
mode="bilinear",
align_corners=False,
).unsqueeze(0)
videos.append(history_video_k)
videos = torch.cat(videos, dim=0) / 255.0 * 2.0 - 1.0
vae_device = next(streaming_vae.vae.parameters()).device
videos_chunk = videos.to(vae_device).to(DTYPE)
enc_out = streaming_vae.encode_chunk(videos_chunk)
mu, logvar = torch.chunk(enc_out, 2, dim=1)
latents_mean = torch.tensor(vae.config.latents_mean).to(mu.device)
latents_std = torch.tensor(vae.config.latents_std).to(mu.device)
mu_norm = normalize_latents(mu, latents_mean, 1.0 / latents_std)
video_latent = torch.cat(mu_norm.split(1, dim=0), dim=-1)
return video_latent.to(DEVICE)
def _repeat_input_for_cfg(input_dict, use_cfg, prompt_embeds, negative_prompt_embeds):
if use_cfg:
input_dict["noisy_latents"] = input_dict["noisy_latents"].repeat(2, 1, 1, 1, 1)
input_dict["text_emb"] = torch.cat(
[prompt_embeds.to(DTYPE).clone(), negative_prompt_embeds.to(DTYPE).clone()],
dim=0,
)
input_dict["grid_id"] = input_dict["grid_id"][None].repeat(2, 1, 1)
input_dict["timesteps"] = input_dict["timesteps"][None].repeat(2, 1)
else:
input_dict["grid_id"] = input_dict["grid_id"][None]
input_dict["timesteps"] = input_dict["timesteps"][None]
return input_dict
def _prepare_latent_input(
latent_model_input,
action_model_input,
latent_t=0,
action_t=0,
latent_cond=None,
action_cond=None,
frame_st_id=0,
patch_size=(1, 2, 2),
prompt_embeds=None,
use_cfg=False,
negative_prompt_embeds=None,
):
input_dict = dict()
if latent_model_input is not None:
input_dict["latent_res_lst"] = {
"noisy_latents": latent_model_input,
"timesteps": torch.ones(
[latent_model_input.shape[2]], dtype=torch.float32, device=DEVICE
) * latent_t,
"grid_id": get_mesh_id(
latent_model_input.shape[-3] // patch_size[0],
latent_model_input.shape[-2] // patch_size[1],
latent_model_input.shape[-1] // patch_size[2],
0,
1,
frame_st_id,
).to(DEVICE),
"text_emb": prompt_embeds.to(DTYPE).clone(),
}
if latent_cond is not None:
input_dict["latent_res_lst"]["noisy_latents"][:, :, 0:1] = latent_cond[:, :, 0:1]
input_dict["latent_res_lst"]["timesteps"][0:1] *= 0
if action_model_input is not None:
input_dict["action_res_lst"] = {
"noisy_latents": action_model_input,
"timesteps": torch.ones(
[action_model_input.shape[2]], dtype=torch.float32, device=DEVICE
) * action_t,
"grid_id": get_mesh_id(
action_model_input.shape[-3],
action_model_input.shape[-2],
action_model_input.shape[-1],
1,
1,
frame_st_id,
action=True,
).to(DEVICE),
"text_emb": prompt_embeds.to(DTYPE).clone(),
}
if action_cond is not None:
input_dict["action_res_lst"]["noisy_latents"][:, :, 0:1] = action_cond[:, :, 0:1]
input_dict["action_res_lst"]["timesteps"][0:1] *= 0
input_dict["action_res_lst"]["noisy_latents"][:, ~action_mask] *= 0
return input_dict
def infer_chunk(
init_latent,
frame_st_id,
prompt_embeds,
negative_prompt_embeds,
use_cfg,
guidance_scale,
action_guidance_scale,
num_chunks_to_infer,
):
"""Generate one video chunk (video latents + action latents)."""
frame_chunk_size = CONFIG["frame_chunk_size"]
latent_height = CONFIG["height"] // 16
latent_width = (CONFIG["width"] // 16) * len(CONFIG["obs_cam_keys"])
latents = torch.randn(
1, 48, frame_chunk_size, latent_height, latent_width, device=DEVICE, dtype=DTYPE
)
actions = torch.randn(
1,
CONFIG["action_dim"],
frame_chunk_size,
CONFIG["action_per_frame"],
1,
device=DEVICE,
dtype=DTYPE,
)
video_inference_step = CONFIG["num_inference_steps"]
action_inference_step = CONFIG["action_num_inference_steps"]
video_step = CONFIG["video_exec_step"]
scheduler.set_timesteps(video_inference_step)
action_scheduler.set_timesteps(action_inference_step)
timesteps = scheduler.timesteps
action_timesteps = action_scheduler.timesteps
timesteps = F.pad(timesteps, (0, 1), mode="constant", value=0)
if video_step != -1:
timesteps = timesteps[:video_step]
action_timesteps = F.pad(action_timesteps, (0, 1), mode="constant", value=0)
with torch.no_grad():
# 1. Video generation loop
for i, t in enumerate(timesteps):
last_step = i == len(timesteps) - 1
latent_cond = init_latent[:, :, 0:1].to(DTYPE) if frame_st_id == 0 else None
input_dict = _prepare_latent_input(
latents,
None,
t,
t,
latent_cond,
None,
frame_st_id=frame_st_id,
patch_size=CONFIG["patch_size"],
prompt_embeds=prompt_embeds,
use_cfg=use_cfg,
negative_prompt_embeds=negative_prompt_embeds,
)
video_noise_pred = transformer(
_repeat_input_for_cfg(
input_dict["latent_res_lst"],
use_cfg,
prompt_embeds,
negative_prompt_embeds,
),
update_cache=1 if last_step else 0,
cache_name="pos",
action_mode=False,
)
if not last_step or video_step != -1:
video_noise_pred = data_seq_to_patch(
CONFIG["patch_size"],
video_noise_pred,
frame_chunk_size,
latent_height,
latent_width,
batch_size=2 if use_cfg else 1,
)
if guidance_scale > 1:
video_noise_pred = video_noise_pred[1:] + guidance_scale * (
video_noise_pred[:1] - video_noise_pred[1:]
)
else:
video_noise_pred = video_noise_pred[:1]
latents = scheduler.step(video_noise_pred, t, latents, return_dict=False)
latents[:, :, 0:1] = (
latent_cond if frame_st_id == 0 else latents[:, :, 0:1]
)
# 2. Action generation loop
for i, t in enumerate(action_timesteps):
last_step = i == len(action_timesteps) - 1
action_cond = (
torch.zeros(
[1, CONFIG["action_dim"], 1, CONFIG["action_per_frame"], 1],
device=DEVICE,
dtype=DTYPE,
)
if frame_st_id == 0
else None
)
input_dict = _prepare_latent_input(
None,
actions,
t,
t,
None,
action_cond,
frame_st_id=frame_st_id,
patch_size=CONFIG["patch_size"],
prompt_embeds=prompt_embeds,
use_cfg=use_cfg,
negative_prompt_embeds=negative_prompt_embeds,
)
action_noise_pred = transformer(
_repeat_input_for_cfg(
input_dict["action_res_lst"],
use_cfg,
prompt_embeds,
negative_prompt_embeds,
),
update_cache=1 if last_step else 0,
cache_name="pos",
action_mode=True,
)
if not last_step:
action_noise_pred = rearrange(
action_noise_pred, "b (f n) c -> b c f n 1", f=frame_chunk_size
)
if action_guidance_scale > 1:
action_noise_pred = action_noise_pred[1:] + action_guidance_scale * (
action_noise_pred[:1] - action_noise_pred[1:]
)
else:
action_noise_pred = action_noise_pred[:1]
actions = action_scheduler.step(
action_noise_pred, t, actions, return_dict=False
)
actions[:, :, 0:1] = (
action_cond if frame_st_id == 0 else actions[:, :, 0:1]
)
actions[:, ~action_mask] *= 0
return actions, latents
def decode_video(pred_latent):
"""Decode latent tensor to video frames."""
vae_device = next(vae.parameters()).device
latents = pred_latent.to(vae_device).to(vae.dtype)
latents_mean = (
torch.tensor(vae.config.latents_mean)
.view(1, vae.config.z_dim, 1, 1, 1)
.to(latents.device, latents.dtype)
)
latents_std = 1.0 / torch.tensor(vae.config.latents_std).view(
1, vae.config.z_dim, 1, 1, 1
).to(latents.device, latents.dtype)
latents = latents / latents_std + latents_mean
with torch.no_grad():
video = vae.decode(latents, return_dict=False)[0]
video = video_processor.postprocess_video(video, output_type="np")[0]
return video
# ------------------------------------------------------------------ #
# Gradio inference function
# ------------------------------------------------------------------ #
@spaces.GPU(duration=60, size="xlarge")
def generate(
top_img: "np.ndarray",
wrist_img: "np.ndarray",
prompt: str,
num_chunks: int = 5,
seed: int = 0,
progress=gr.Progress(track_tqdm=True),
):
"""Generate a robot manipulation video from initial observations and a text prompt.
Args:
top_img: Top-down camera observation image.
wrist_img: Wrist camera observation image.
prompt: Natural language instruction for the robot task.
num_chunks: Number of video chunks to generate autoregressively (each chunk = 4 frames).
seed: Random seed for reproducibility.
"""
torch.manual_seed(seed)
torch.cuda.manual_seed(seed)
use_cfg = CONFIG["guidance_scale"] > 1 or CONFIG["action_guidance_scale"] > 1
# Prepare observations
obs = [
{
CONFIG["obs_cam_keys"][0]: top_img,
CONFIG["obs_cam_keys"][1]: wrist_img,
}
]
# Reset KV cache
transformer.clear_cache("pos")
streaming_vae.clear_cache()
# Encode initial observation
init_latent = encode_obs(obs)
# Encode prompt
prompt_embeds, negative_prompt_embeds = encode_prompt(prompt)
# Latent dimensions
latent_height = CONFIG["height"] // 16
latent_width = (CONFIG["width"] // 16) * len(CONFIG["obs_cam_keys"])
patch_size = CONFIG["patch_size"]
latent_token_per_chunk = (
CONFIG["frame_chunk_size"] * latent_height * latent_width
) // (patch_size[0] * patch_size[1] * patch_size[2])
action_token_per_chunk = CONFIG["frame_chunk_size"] * CONFIG["action_per_frame"]
# Create KV cache
transformer.create_empty_cache(
"pos",
CONFIG["attn_window"],
latent_token_per_chunk,
action_token_per_chunk,
device=DEVICE,
dtype=DTYPE,
batch_size=2 if use_cfg else 1,
)
# Autoregressive chunk generation
pred_latent_lst = []
for chunk_id in range(num_chunks):
frame_st_id = chunk_id * CONFIG["frame_chunk_size"]
actions, latents = infer_chunk(
init_latent,
frame_st_id,
prompt_embeds,
negative_prompt_embeds,
use_cfg,
CONFIG["guidance_scale"],
CONFIG["action_guidance_scale"],
num_chunks,
)
pred_latent_lst.append(latents)
pred_latent = torch.cat(pred_latent_lst, dim=2)
# Free VRAM before VAE decode: move transformer and text encoder to CPU
transformer.clear_cache("pos")
streaming_vae.clear_cache()
transformer.to("cpu")
text_encoder.to("cpu")
torch.cuda.empty_cache()
# Decode video
pred_latent_cpu = pred_latent.cpu()
del pred_latent
torch.cuda.empty_cache()
video = decode_video(pred_latent_cpu)
# Save to temp file
from diffusers.utils import export_to_video
tmp_file = tempfile.NamedTemporaryFile(suffix=".mp4", delete=False)
tmp_file.close()
export_to_video(video, tmp_file.name, fps=10)
return tmp_file.name
# ------------------------------------------------------------------ #
# Gradio UI
# ------------------------------------------------------------------ #
CSS = """
#col-container { max-width: 1100px; margin: 0 auto; }
.dark .gradio-container { color: var(--body-text-color); }
"""
with gr.Blocks() as demo:
gr.Markdown(
"""
# Next Forcing: Causal World Modeling with Multi-Chunk Prediction
Generate robot manipulation video from initial observations and a text instruction.
Upload top-down and wrist camera images, describe the task, and the model autoregressively
predicts future video frames.
[Paper](https://arxiv.org/abs/2606.11187) | [Code](https://github.com/gangweix/next-forcing) | [Model](https://huggingface.co/gangweix/next-forcing-base)
"""
)
with gr.Row():
with gr.Column():
top_img = gr.Image(
label="Top Camera",
type="numpy",
height=256,
)
wrist_img = gr.Image(
label="Wrist Camera",
type="numpy",
height=256,
)
prompt = gr.Textbox(
label="Task Instruction",
placeholder="e.g. Pick the green cube and place it inside the blue box",
lines=2,
)
with gr.Accordion("Advanced Settings", open=False):
num_chunks = gr.Slider(
label="Number of chunks (4 frames each)",
minimum=1,
maximum=10,
value=5,
step=1,
)
seed = gr.Number(label="Seed", value=0, precision=0)
run_btn = gr.Button("Generate Video", variant="primary")
with gr.Column():
video_out = gr.Video(label="Generated Video")
gr.Examples(
examples=[
[
"examples/observation.images.top.png",
"examples/observation.images.wrist.png",
"Pick the green cube and place it inside the blue box",
5,
0,
],
[
"examples/observation.images.top.png",
"examples/observation.images.wrist.png",
"Move the red block to the left side of the table",
5,
42,
],
],
inputs=[top_img, wrist_img, prompt, num_chunks, seed],
outputs=video_out,
fn=generate,
cache_examples=True,
cache_mode="lazy",
)
run_btn.click(
fn=generate,
inputs=[top_img, wrist_img, prompt, num_chunks, seed],
outputs=video_out,
api_name="generate",
)
demo.launch(mcp_server=True, theme=gr.themes.Citrus(), css=CSS)