import os os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True") import spaces # must be imported before torch import random import time import uuid import numpy as np import torch from huggingface_hub import hf_hub_download, snapshot_download from omegaconf import OmegaConf # --------------------------------------------------------------------------- # Download base-model assets (T5 text encoder, VAE, tokenizer) and the SGF+ # checkpoint. The base DiT safetensors is NOT needed: the SGF+ checkpoint # contains the full dual-expert generator (generation + memory parameters). # --------------------------------------------------------------------------- snapshot_download( repo_id="Wan-AI/Wan2.1-T2V-1.3B", local_dir="wan_models/Wan2.1-T2V-1.3B", allow_patterns=[ "models_t5_umt5-xxl-enc-bf16.pth", "Wan2.1_VAE.pth", "google/umt5-xxl/*", ], ) hf_hub_download( repo_id="ZihanSu/Self_Gradient_Forcing_Plus", filename="chunkwise/model.pt", local_dir=".", ) import av import gradio as gr import tqdm as tqdm_pkg # noqa: F401 (pipeline/causal_inference.py imports tqdm) from demo_utils.constant import ZERO_VAE_CACHE from demo_utils.vae_block3 import VAEDecoderWrapper from pipeline import CausalInferencePipeline from utils.wan_wrapper import WanDiffusionWrapper, WanTextEncoder from utils.scheduler import FlowMatchScheduler # noqa: F401 (imported by wrapper) from wan.modules.causal_model import CausalWanModel DEVICE = "cuda" WAN_DIR = "wan_models/Wan2.1-T2V-1.3B" CKPT_PATH = "chunkwise/model.pt" NUM_LATENT_FRAMES = 21 # 21 latents -> 81 pixel frames @ 480x832, matching the paper's clip length FPS = 16 class SGFDiffusionWrapper(WanDiffusionWrapper): """WanDiffusionWrapper that reuses an already-constructed CausalWanModel. The upstream wrapper's __init__ calls CausalWanModel.from_pretrained, which would download the 5.7 GB base diffusion_pytorch_model.safetensors even though the SGF+ checkpoint below overwrites every generator weight (both the generation and the memory expert). Building the model directly from the 1.3B architecture (same params as wan_models/Wan2.1-T2V-1.3B/config.json) and injecting it here skips that download while keeping the wrapper's forward/scheduler contract intact. """ def __init__(self, model: CausalWanModel, timestep_shift: float = 5.0): torch.nn.Module.__init__(self) self.model = model self.model.eval() # For non-causal diffusion, all frames share the same timestep self.uniform_timestep = False self.scheduler = FlowMatchScheduler( shift=timestep_shift, sigma_min=0.0, extra_one_step=True ) self.scheduler.set_timesteps(1000, training=True) self.seq_len = 32760 # [1, 21, 16, 60, 104] self.post_init() # --------------------------------------------------------------------------- # Config: authors' sgf_plus_chunkwise.yaml merged over default_config.yaml, # exactly like inference.py does. # --------------------------------------------------------------------------- config = OmegaConf.load("configs/sgf_plus_chunkwise.yaml") default_config = OmegaConf.load("configs/default_config.yaml") config = OmegaConf.merge(default_config, config) print("Initializing models...", flush=True) # --- Text encoder (UMT5-XXL) --- text_encoder = WanTextEncoder() # --- Dual-expert generator --- # Construct the causal Wan DiT directly from the 1.3B architecture (same # params as wan_models/Wan2.1-T2V-1.3B/config.json) instead of loading the # base diffusion_pytorch_model.safetensors — the SGF+ checkpoint below # overwrites every generator weight anyway (it contains both the generation # and the memory expert). model_kwargs = config.model_kwargs dit = CausalWanModel( model_type="t2v", patch_size=(1, 2, 2), text_len=512, in_dim=16, dim=1536, ffn_dim=8960, freq_dim=256, text_dim=4096, out_dim=16, num_heads=12, num_layers=30, local_attn_size=model_kwargs.local_attn_size, sink_size=model_kwargs.sink_size, ) dit.enable_dual_full_expert() transformer = SGFDiffusionWrapper(dit, timestep_shift=model_kwargs.timestep_shift) state_dict = torch.load(CKPT_PATH, map_location="cpu", weights_only=False) gen_sd = state_dict.get("generator_ema", state_dict.get("generator")) try: transformer.load_state_dict(gen_sd, strict=True) except RuntimeError: fixed = {} for k, v in gen_sd.items(): if k.startswith("model._fsdp_wrapped_module."): k = k.replace("model._fsdp_wrapped_module.", "", 1) fixed[k] = v transformer.load_state_dict(fixed, strict=True) del state_dict, gen_sd # --- VAE decoder (block-cached, used for streaming decode) --- vae_decoder = VAEDecoderWrapper() vae_state_dict = torch.load(f"{WAN_DIR}/Wan2.1_VAE.pth", map_location="cpu") decoder_state_dict = { k: v for k, v in vae_state_dict.items() if "decoder." in k or "conv2" in k } vae_decoder.load_state_dict(decoder_state_dict) del vae_state_dict, decoder_state_dict text_encoder.eval().to(dtype=torch.bfloat16).requires_grad_(False).to(DEVICE) transformer.eval().to(dtype=torch.float16).requires_grad_(False).to(DEVICE) vae_decoder.eval().to(dtype=torch.float16).requires_grad_(False).to(DEVICE) pipeline = CausalInferencePipeline( config, device=DEVICE, generator=transformer, text_encoder=text_encoder, vae=vae_decoder, ) # NOTE: no pipeline-level dtype cast — the UMT5 text encoder must stay in # bfloat16 (that is the dtype of its released weights); the DiT and VAE # decoder were cast to float16 individually above. pipeline.to(DEVICE) print("Models ready.", flush=True) def frames_to_ts_file(frames: list, filepath: str, fps: int) -> str: """Encode a list of HWC RGB uint8 frames into an MPEG-TS chunk with PyAV.""" height, width = frames[0].shape[:2] container = av.open(filepath, mode="w", format="mpegts") stream = container.add_stream("h264", rate=fps) stream.width = width stream.height = height stream.pix_fmt = "yuv420p" stream.options = { "preset": "ultrafast", "tune": "zerolatency", "crf": "23", "profile": "baseline", "level": "3.0", } try: for frame_np in frames: frame = av.VideoFrame.from_ndarray(frame_np, format="rgb24") frame = frame.reformat(format=stream.pix_fmt) for packet in stream.encode(frame): container.mux(packet) for packet in stream.encode(): container.mux(packet) finally: container.close() return filepath @torch.no_grad() # duration sized to the measured worst case: a full 7-block / 81-frame clip costs # ~21 s wall on a cold zero-a10g worker (~18 s of pure denoise + decode), so 30 s # = round(21 * 1.4) per the ZeroGPU sizing guidance. Leaner durations rank higher # in the queue and spend less of each visitor's daily quota. @spaces.GPU(duration=30) def generate( prompt: str, seed: int = -1, num_blocks: int = 7, fps: int = FPS, ): """Generate a short video from a text prompt with SGF+ (chunkwise). Streams the video block-by-block as MPEG-TS chunks: each latent block (3 latents = 12 pixel frames) is denoised in 4 few-step iterations, written to the KV cache through the memory expert, decoded with the block-cached VAE, and yielded as soon as it is ready. Args: prompt: Text prompt describing the video to generate. seed: RNG seed for the initial noise (-1 for random). num_blocks: Number of 3-latent blocks to generate (7 => 81 frames). fps: Playback frames-per-second of the output stream. Yields: (video_chunk_path, seed, status_html) tuples; a chunk path of None leaves the streamed video untouched while updating the status. """ if seed is None or int(seed) < 0: seed = random.randint(0, 2**31 - 1) seed = int(seed) t0 = time.perf_counter() os.makedirs("gradio_tmp", exist_ok=True) conditional_dict = text_encoder(text_prompts=[prompt]) conditional_dict = { k: v.to(dtype=torch.float16) for k, v in conditional_dict.items() } rnd = torch.Generator(DEVICE).manual_seed(seed) noise = torch.randn( [1, NUM_LATENT_FRAMES, 16, 60, 104], device=DEVICE, dtype=torch.float16, generator=rnd, ) pipeline._initialize_kv_cache(1, torch.float16, DEVICE) pipeline._initialize_crossattn_cache(1, torch.float16, DEVICE) vae_cache = [c.to(device=DEVICE, dtype=torch.float16) for c in ZERO_VAE_CACHE] num_blocks = max(1, int(num_blocks)) all_num_frames = [pipeline.num_frame_per_block] * num_blocks current_start_frame = 0 total_frames = 0 elapsed = None for idx, current_num_frames in enumerate(all_num_frames): noisy_input = noise[ :, current_start_frame:current_start_frame + current_num_frames ] # --- 4-step few-step denoising (generation expert) --- for step_idx, current_timestep in enumerate(pipeline.denoising_step_list): timestep = ( torch.ones([1, current_num_frames], device=noise.device, dtype=torch.int64) * current_timestep ) _, denoised_pred = pipeline.generator( noisy_image_or_video=noisy_input, conditional_dict=conditional_dict, timestep=timestep, kv_cache=pipeline.kv_cache1, crossattn_cache=pipeline.crossattn_cache, current_start=current_start_frame * pipeline.frame_seq_length, use_memory_expert=False, ) if step_idx < len(pipeline.denoising_step_list) - 1: next_timestep = pipeline.denoising_step_list[step_idx + 1] noisy_input = pipeline.scheduler.add_noise( denoised_pred.flatten(0, 1), torch.randn_like(denoised_pred.flatten(0, 1)), next_timestep * torch.ones( [1 * current_num_frames], device=noise.device, dtype=torch.long, ), ).unflatten(0, denoised_pred.shape[:2]) # --- Write clean context to the KV cache (memory expert) --- pipeline.generator( noisy_image_or_video=denoised_pred, conditional_dict=conditional_dict, timestep=torch.ones_like(timestep) * config.context_noise, kv_cache=pipeline.kv_cache1, crossattn_cache=pipeline.crossattn_cache, current_start=current_start_frame * pipeline.frame_seq_length, use_memory_expert=True, ) # --- Decode this block to pixels with the cached VAE --- pixels, vae_cache = pipeline.vae(denoised_pred.half(), *vae_cache) if idx == 0: pixels = pixels[:, 3:] # skip latent-0 duplication of first 3 frames frames = [] for f in range(pixels.shape[1]): frame = torch.clamp(pixels[0, f].float(), -1.0, 1.0) * 127.5 + 127.5 frames.append(frame.to(torch.uint8).cpu().numpy().transpose(1, 2, 0)) total_frames += len(frames) ts_path = os.path.join("gradio_tmp", f"block_{idx:04d}_{uuid.uuid4().hex[:8]}.ts") frames_to_ts_file(frames, ts_path, fps) progress = min(100.0, 100.0 * (idx + 1) / num_blocks) status = ( f"