import os os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True") import sys import time import tempfile import spaces # noqa: E402 (must precede torch / CUDA imports) import torch import numpy as np import gradio as gr from PIL import Image from huggingface_hub import hf_hub_download # Make the vendored diffsynth package importable. SCRIPT_DIR = os.path.dirname(os.path.abspath(__file__)) if SCRIPT_DIR not in sys.path: sys.path.insert(0, SCRIPT_DIR) from einops import rearrange from diffsynth.models.utils import load_state_dict from diffsynth.models.wan_video_dit import sinusoidal_embedding_1d from diffsynth.models.wan_video_dit_dual_stream import init_flow_stream from diffsynth.pipelines.wan_video_new import WanVideoPipeline, ModelConfig from diffsynth.pipelines.wan_video_dual_stream import _dual_stream_block_fn from diffsynth.data.video import save_video # ---------------------------------------------------------------------------- # Model setup (module scope — ZeroGPU packs weights to disk at startup). # ---------------------------------------------------------------------------- BASE_MODEL = "Wan-AI/Wan2.2-TI2V-5B" TOKENIZER_MODEL = "Wan-AI/Wan2.1-T2V-1.3B" FLOWWAM_REPO = "YixiangChen/FlowWAM" FLOWWAM_CKPT = "flowwam_worldarena_stage1.safetensors" # This Space loads the FlowWAM *WorldArena world-model* checkpoint # (flowwam_worldarena_stage1.safetensors). Per the paper, in WORLD-MODEL MODE # the flow stream is NOT denoised: the flow latents are set to the clean VAE # encoding of a desired motion trajectory and held FIXED throughout sampling, # while only the RGB latents are initialised from noise and denoised. The # model is conditioned on the initial frame + a language instruction (no # RoboTwin T-shape camera prefix — that belongs to the separate # flowwam_robotwin action checkpoint, and prepending it drives the # world-model checkpoint off-distribution). With no external flow input in a # generic image+text demo, the "desired motion" is a static (zero-motion) # field: a fully-white flow video (the FlowCodec zero-flow sentinel). MODELS_DIR = os.path.join(SCRIPT_DIR, "models") os.makedirs(MODELS_DIR, exist_ok=True) DTYPE = torch.bfloat16 DEVICE = "cuda" def _mc(pattern, offload="cpu"): return ModelConfig( model_id=BASE_MODEL, origin_file_pattern=pattern, offload_device=offload, local_model_path=MODELS_DIR, download_resource="huggingface", ) print("Loading Wan2.2-TI2V-5B dual-stream pipeline (VAE + T5 + DiT) ...", flush=True) pipe = WanVideoPipeline.from_pretrained( torch_dtype=DTYPE, device=DEVICE, model_configs=[ _mc("models_t5_umt5-xxl-enc-bf16.pth"), _mc("diffusion_pytorch_model*.safetensors"), _mc("Wan2.2_VAE.pth"), ], tokenizer_config=ModelConfig( model_id=TOKENIZER_MODEL, origin_file_pattern="google/*", local_model_path=MODELS_DIR, download_resource="huggingface", ), redirect_common_files=False, ) # Flow stream: deep-copied patch-embed + head from the DiT. flow_stream = init_flow_stream(pipe.dit) # Load the FlowWAM checkpoint: DiT + flow_stream keys (no action_expert in # the world-model stage-1 checkpoint). print(f"Downloading FlowWAM checkpoint {FLOWWAM_CKPT} ...", flush=True) ckpt_path = hf_hub_download(FLOWWAM_REPO, FLOWWAM_CKPT) state_dict = load_state_dict(ckpt_path) dit_keys, flow_keys = {}, {} for k, v in state_dict.items(): if k.startswith("action_expert."): continue if k.startswith("flow_stream."): flow_keys[k.replace("flow_stream.", "")] = v else: dit_keys[k] = v # Params trained in fp32 (modulation / time-MLP / LayerNorm) — restore later. fp32_dit_values = {k: v.clone() for k, v in dit_keys.items() if v.dtype == torch.float32} missing, unexpected = pipe.dit.load_state_dict(dit_keys, strict=False) print(f"DiT (full): loaded {len(dit_keys) - len(unexpected)} keys, " f"{len(missing)} missing, {len(unexpected)} unexpected", flush=True) missing, unexpected = flow_stream.load_state_dict(flow_keys, strict=False) print(f"FlowStream (full): loaded {len(flow_keys) - len(unexpected)} keys, " f"{len(missing)} missing, {len(unexpected)} unexpected", flush=True) pipe.enable_vram_management() def _apply_fp32_modulation(dit, fp32_state_values): """Restore fp32 precision for modulation / time-MLP / LayerNorm params.""" from diffsynth.vram_management.layers import AutoWrappedLinear, WanAutoCastLayerNorm param_map = dict(dit.named_parameters()) for key, fp32_value in fp32_state_values.items(): if key in param_map: param_map[key].data = fp32_value.to(device=param_map[key].device) for seq_module in [dit.time_embedding, dit.time_projection]: for sub in seq_module.modules(): if isinstance(sub, AutoWrappedLinear): sub.offload_dtype = torch.float32 sub.onload_dtype = torch.float32 sub.computation_dtype = torch.float32 def _pre_hook(_mod, args): return tuple(a.float() if isinstance(a, torch.Tensor) else a for a in args) def _post_hook(_mod, _args, output): return output.bfloat16() if isinstance(output, torch.Tensor) else output for seq_module in [dit.time_embedding, dit.time_projection]: seq_module.register_forward_pre_hook(_pre_hook) seq_module.register_forward_hook(_post_hook) for module in dit.modules(): if isinstance(module, WanAutoCastLayerNorm): module.offload_dtype = torch.float32 module.onload_dtype = torch.float32 if fp32_dit_values: _apply_fp32_modulation(pipe.dit, fp32_dit_values) flow_stream = flow_stream.to(device=DEVICE, dtype=DTYPE).eval() print("FlowWAM pipeline ready.", flush=True) # ---------------------------------------------------------------------------- # World-model forward: RGB is denoised at the sampling timestep while the flow # stream is held FIXED at its clean VAE latent (per the FlowWAM paper's # world-model mode). This reuses the exact dual-stream block math # (``_dual_stream_block_fn``) but labels the clean flow tokens with timestep 0 # (like the reference's clean video-conditioning pass), instead of tying both # streams to the same noisy timestep. Only ``rgb_out`` is used. # ---------------------------------------------------------------------------- @torch.no_grad() def _world_model_rgb_pred(dit, flow_stream, rgb_latents, flow_clean_latents, rgb_timestep, context): B = rgb_latents.shape[0] dtype = rgb_latents.dtype dev = rgb_latents.device # Per-token timestep: RGB first frame = 0 (I2V prefix), other RGB frames = # ts_b; ALL flow tokens = 0 because the flow stream is clean and fixed. rgb_spatial = rgb_latents.shape[3] * rgb_latents.shape[4] // 4 rgb_temporal = rgb_latents.shape[2] flow_spatial = flow_clean_latents.shape[3] * flow_clean_latents.shape[4] // 4 flow_temporal = flow_clean_latents.shape[2] t_per_token_list = [] for b in range(B): ts_b = (rgb_timestep[b] if rgb_timestep.dim() >= 1 and rgb_timestep.shape[0] > 1 else rgb_timestep) rgb_tpt = torch.cat([ torch.zeros(1, rgb_spatial, dtype=dtype, device=dev), torch.ones(rgb_temporal - 1, rgb_spatial, dtype=dtype, device=dev) * ts_b, ]).flatten() # Flow stream is clean everywhere -> timestep 0 for every flow token. flow_tpt = torch.zeros(flow_temporal * flow_spatial, dtype=dtype, device=dev) t_per_token_list.append(torch.cat([rgb_tpt, flow_tpt])) t_per_token = torch.stack(t_per_token_list, dim=0) t = dit.time_embedding( sinusoidal_embedding_1d(dit.freq_dim, t_per_token.reshape(-1)) .reshape(B, -1, dit.freq_dim) ) t_mod = dit.time_projection(t).unflatten(2, (6, dit.dim)) context = dit.text_embedding(context) rgb_5d = dit.patchify(rgb_latents) f_r, h_r, w_r = rgb_5d.shape[2:] rgb_tokens = rearrange(rgb_5d, 'b c f h w -> b (f h w) c').contiguous() n_rgb = rgb_tokens.shape[1] n_rgb_tok = rgb_spatial * rgb_temporal t_rgb = t[:, :n_rgb_tok] flow_5d = flow_stream.patchify(flow_clean_latents) f_f, h_f, w_f = flow_5d.shape[2:] flow_tokens = rearrange(flow_5d, 'b c f h w -> b (f h w) c').contiguous() flow_tokens = flow_tokens + flow_stream.stream_embed.to(dtype=flow_tokens.dtype, device=flow_tokens.device) rgb_freqs = torch.cat([ dit.freqs[0][:f_r].view(f_r, 1, 1, -1).expand(f_r, h_r, w_r, -1), dit.freqs[1][:h_r].view(1, h_r, 1, -1).expand(f_r, h_r, w_r, -1), dit.freqs[2][:w_r].view(1, 1, w_r, -1).expand(f_r, h_r, w_r, -1), ], dim=-1).reshape(f_r * h_r * w_r, 1, -1).to(rgb_tokens.device) flow_freqs = torch.cat([ dit.freqs[0][:f_f].view(f_f, 1, 1, -1).expand(f_f, h_f, w_f, -1), dit.freqs[1][:h_f].view(1, h_f, 1, -1).expand(f_f, h_f, w_f, -1), dit.freqs[2][:w_f].view(1, 1, w_f, -1).expand(f_f, h_f, w_f, -1), ], dim=-1).reshape(f_f * h_f * w_f, 1, -1).to(flow_tokens.device) for block in dit.blocks: rgb_tokens, flow_tokens = _dual_stream_block_fn( block, rgb_tokens, flow_tokens, context, t_mod, rgb_freqs, flow_freqs, n_rgb, ) rgb_out = dit.head(rgb_tokens, t_rgb) rgb_out = dit.unpatchify(rgb_out, (f_r, h_r, w_r)) return rgb_out # ---------------------------------------------------------------------------- # Inference — dual-stream world-model rollout (stage 1 only). # ---------------------------------------------------------------------------- def _estimate(image, instruction, num_frames=49, num_inference_steps=25, sigma_shift=5.0, seed=1, *args, **kwargs): # Measured: ~38s warm for 49 frames / 25 steps; cold start adds ~15-20s. steps = int(num_inference_steps) return min(120, 35 + int(steps * 2.2)) @spaces.GPU(duration=_estimate) @torch.no_grad() def generate(image, instruction, num_frames=49, num_inference_steps=25, sigma_shift=5.0, seed=1, progress=gr.Progress(track_tqdm=True)): """Generate a future RGB video from one image + instruction (world-model mode). Runs FlowWAM's WorldArena world-model checkpoint in flow-conditioned mode: the flow stream is held fixed at the clean encoding of a (static) motion trajectory and only the RGB stream is denoised, conditioned on the first frame and the instruction. Args: image: the conditioning first frame (PIL image). instruction: text describing the action / motion to imagine. num_frames: number of frames to generate (4k+1). num_inference_steps: RGB denoising steps. sigma_shift: flow-match scheduler sigma shift. seed: RNG seed. Returns: (rgb_video_path, flow_video_path): mp4 files for the generated future RGB frames and the fixed flow-conditioning trajectory. """ if image is None: raise gr.Error("Please provide an input image.") instruction = (instruction or "").strip() device = pipe.device dtype = pipe.torch_dtype vae_z_dim = getattr(pipe.vae, "z_dim", 16) seed = int(seed) num_frames = int(num_frames) # ---- Resize conditioning frame to a valid grid ---- if isinstance(image, np.ndarray): image = Image.fromarray(image) image = image.convert("RGB") w, h = image.size # Keep a compact aspect-preserving size (~320x256 like the reference). target_w = 320 target_h = max(1, round(h * target_w / w)) tiled_h, tiled_w, video_frames = pipe.check_resize_height_width( target_h, target_w, num_frames) cond_pil = image.resize((tiled_w, tiled_h), Image.BICUBIC) # ---- Text encoding ---- # World-model conditioning is the initial frame + the plain language # instruction (no RoboTwin camera prefix — see the note above). pipe.load_models_to_device(["text_encoder"]) context = pipe.prompter.encode_prompt(instruction, positive=True, device=device) # ---- VAE encode conditioning frame + fixed clean flow trajectory ---- pipe.load_models_to_device(["vae"]) upscale = pipe.vae.upsampling_factor T_lat = (video_frames - 1) // 4 + 1 rgb_H_lat = tiled_h // upscale rgb_W_lat = tiled_w // upscale # RGB: clean latent of the first frame (fixed as the I2V prefix). rgb_vid = pipe.preprocess_video([cond_pil]) rgb_prefix = pipe.vae.encode(rgb_vid, device=device).to(dtype=dtype, device=device) # Flow: WORLD-MODEL MODE — the flow latents are the clean VAE encoding of # the desired motion trajectory, held fixed throughout sampling. With no # external flow input we use a static (zero-motion) trajectory: a full # white flow video (the FlowCodec zero-flow sentinel). Encode ALL frames # so the entire flow stream is a valid clean latent (not just a prefix). zero_flow_pil = Image.new("RGB", (tiled_w, tiled_h), (255, 255, 255)) flow_vid = pipe.preprocess_video([zero_flow_pil] * video_frames) flow_clean = pipe.vae.encode(flow_vid, device=device).to(dtype=dtype, device=device) rgb_noise_shape = (1, vae_z_dim, T_lat, rgb_H_lat, rgb_W_lat) rgb_noise = pipe.generate_noise(rgb_noise_shape, seed=seed, rand_device="cpu").to(dtype=dtype, device=device) rgb_noise[:, :, :1] = rgb_prefix rgb_latents = rgb_noise.clone() # Flow stream stays clean & fixed for the whole rollout. flow_latents = flow_clean.clone() # ---- World-model video denoising: only RGB is denoised ---- pipe.scheduler.set_timesteps(int(num_inference_steps), shift=float(sigma_shift)) pipe.load_models_to_device(pipe.in_iteration_models) for progress_id, timestep in enumerate(pipe.scheduler.timesteps): t_tensor = timestep.unsqueeze(0).to(dtype=dtype, device=device) rgb_pred = _world_model_rgb_pred( dit=pipe.dit, flow_stream=flow_stream, rgb_latents=rgb_latents, flow_clean_latents=flow_latents, rgb_timestep=t_tensor, context=context, ) rgb_latents = pipe.scheduler.step(rgb_pred, pipe.scheduler.timesteps[progress_id], rgb_latents) rgb_latents[:, :, :1] = rgb_prefix # flow_latents intentionally held fixed (clean conditioning). # ---- Decode: RGB is the generated future; flow is the fixed condition ---- pipe.load_models_to_device(["vae"]) rgb_frames = pipe.vae_output_to_video(pipe.vae.decode(rgb_latents, device=device)) flow_frames = pipe.vae_output_to_video(pipe.vae.decode(flow_latents, device=device)) pipe.load_models_to_device([]) rgb_path = tempfile.NamedTemporaryFile(suffix="_rgb.mp4", delete=False).name flow_path = tempfile.NamedTemporaryFile(suffix="_flow.mp4", delete=False).name save_video(rgb_frames, rgb_path, fps=12) save_video(flow_frames, flow_path, fps=12) return rgb_path, flow_path # ---------------------------------------------------------------------------- # UI # ---------------------------------------------------------------------------- CSS = """ #col-container { max-width: 1100px; margin: 0 auto; } .dark .gradio-container { color: var(--body-text-color); } """ DESCRIPTION = """ # FlowWAM — Optical Flow as a Unified Action Representation A dual-stream video diffusion model (built on **Wan2.2-TI2V-5B**) run in **world-model mode**: the optical-flow stream is held fixed as a clean motion condition while the model denoises a **future RGB video** from one image and a short text instruction. From the paper *FlowWAM: Optical Flow as a Unified Action Representation for World Action Models*. Give it a starting frame and describe the motion — it imagines how the scene evolves under the flow condition. [Paper](https://huggingface.co/papers/2607.13017) · [Code](https://github.com/YixiangChen515/FlowWAM) · [Weights](https://huggingface.co/YixiangChen/FlowWAM) """ with gr.Blocks() as demo: with gr.Column(elem_id="col-container"): gr.Markdown(DESCRIPTION) with gr.Row(): with gr.Column(): image = gr.Image(label="Input image (first frame)", type="pil") instruction = gr.Textbox( label="Instruction", placeholder="describe the action, e.g. 'Hold the gray kitchenpot with both arms'", ) run = gr.Button("Generate", variant="primary") with gr.Column(): rgb_out = gr.Video(label="Generated future RGB") flow_out = gr.Video(label="Flow-conditioning trajectory") with gr.Accordion("Advanced settings", open=False): num_frames = gr.Slider(13, 49, value=49, step=4, label="Frames (4k+1)") num_inference_steps = gr.Slider(10, 40, value=25, step=1, label="Denoising steps") sigma_shift = gr.Slider(1.0, 8.0, value=5.0, step=0.5, label="Sigma shift") seed = gr.Number(value=1, precision=0, label="Seed") inputs = [image, instruction, num_frames, num_inference_steps, sigma_shift, seed] run.click(generate, inputs=inputs, outputs=[rgb_out, flow_out], api_name="generate") # Real RoboTwin first-frames + bare task instructions from the # reference dataset (YixiangChen/FlowWAM_RoboTwin, aloha-agilex_clean_50, # head-camera frame 0 of episode0). These match the world-model # checkpoint's training distribution: a 320x240 tabletop aloha-robot # scene + a plain manipulation instruction with NO RoboTwin camera # prefix (the prefix belongs to the separate flowwam_robotwin action # checkpoint and would drive this world-model checkpoint # off-distribution). gr.Examples( examples=[ ["robotwin_lift_pot.png", "Hold the gray kitchenpot with both arms"], ["robotwin_open_laptop.png", "Lift and open the laptop with black textured screen."], ["robotwin_place_bread_basket.png", "Pick up both bread loaves and place them in the white oval breadbasket."], ], inputs=[image, instruction], outputs=[rgb_out, flow_out], fn=generate, cache_examples=True, cache_mode="lazy", ) if __name__ == "__main__": demo.launch(mcp_server=True, theme=gr.themes.Citrus(), css=CSS)