Spaces:
Running on Zero
Running on Zero
Download app.py from hugging-apps/self-gradient-forcing-plus: direct link, hf CLI and curl.
- Browser
- Download file 18 kB
-
https://huggingface.co/spaces/hugging-apps/self-gradient-forcing-plus/resolve/main/app.py
- Command line
-
hf download hf://spaces/hugging-apps/self-gradient-forcing-plus/app.py
-
curl -L -o app.py https://huggingface.co/spaces/hugging-apps/self-gradient-forcing-plus/resolve/main/app.py
18 kB
| 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 | |
| # 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. | |
| 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"<div style='padding:10px;font-family:sans-serif;border:1px solid #ddd;border-radius:8px;'>" | |
| f"<b>Generating…</b> block {idx + 1}/{num_blocks} · {total_frames} frames" | |
| f"<div style='background:#e9ecef;border-radius:4px;margin-top:6px;overflow:hidden;'>" | |
| f"<div style='width:{progress:.0f}%;height:14px;background:#f97316;'></div></div></div>" | |
| ) | |
| yield ts_path, seed, status | |
| current_start_frame += current_num_frames | |
| elapsed = time.perf_counter() - t0 | |
| fps_achieved = total_frames / elapsed if elapsed > 0 else float("inf") | |
| status = ( | |
| f"<div style='padding:12px;font-family:sans-serif;border:1px solid #16a34a;" | |
| f"background:#f0fdf4;border-radius:8px;'>✅ <b>Done</b> — {total_frames} frames " | |
| f"({num_blocks} blocks) in {elapsed:.1f}s · {fps_achieved:.1f} FPS generation</div>" | |
| ) | |
| yield None, seed, status | |
| CSS = """ | |
| .gradio-container, main { max-width: 1100px !important; margin: 0 auto !important; } | |
| .dark .gradio-container { color: var(--body-text-color); } | |
| """ | |
| with gr.Blocks(title="SGF+ Video Generation") as demo: | |
| gr.Markdown( | |
| "# SGF+: Decoupling Gradient Flows for Autoregressive Video Generation\n" | |
| "Few-step autoregressive text-to-video on Wan2.1-T2V-1.3B with dual " | |
| "generation/memory experts. " | |
| "[Paper](https://huggingface.co/papers/2610.10429) · " | |
| "[Code](https://github.com/Zihan-Su/Self_Gradient_Forcing_Plus) · " | |
| "[Model](https://huggingface.co/ZihanSu/Self_Gradient_Forcing_Plus)" | |
| ) | |
| with gr.Row(): | |
| with gr.Column(scale=2): | |
| prompt = gr.Textbox( | |
| label="Prompt", | |
| placeholder="A stylish woman walks down a Tokyo street filled with warm evening neon light...", | |
| lines=4, | |
| ) | |
| with gr.Accordion("Advanced settings", open=False): | |
| seed = gr.Number( | |
| label="Seed (-1 for random)", value=-1, precision=0 | |
| ) | |
| num_blocks = gr.Slider( | |
| label="Blocks (3 latents each)", | |
| minimum=1, maximum=7, value=7, step=1, | |
| info="7 blocks = 81 frames ≈ 5s. Fewer blocks = shorter clip.", | |
| ) | |
| generate_btn = gr.Button("Generate video", variant="primary", size="lg") | |
| with gr.Column(scale=3): | |
| streaming_video = gr.Video( | |
| label="Video stream", streaming=True, autoplay=True, loop=True, height=420 | |
| ) | |
| seed_out = gr.Number(label="Seed used", precision=0, interactive=False) | |
| status_display = gr.HTML( | |
| value=( | |
| "<div style='text-align:center;padding:18px;color:#888;" | |
| "border:1px dashed #ddd;border-radius:8px;'>" | |
| "🎬 Enter a prompt and press <b>Generate video</b>. " | |
| "Frames stream as each block finishes.</div>" | |
| ) | |
| ) | |
| gr.Examples( | |
| examples=[ | |
| ["A stylish woman walks confidently down a bustling Tokyo street at night, neon lights and vibrant city signs glowing around her. She wears a sleek black leather jacket, a flowing red dress, and black boots, with a black purse slung over her shoulder. Her sunglasses rest on her nose, bold red lipstick enhancing her cool, composed expression. The wet pavement mirrors the colorful lights, creating dazzling reflections. Pedestrians blur in the background, adding energy to the scene. Dynamic medium shot, slight side angle, smooth motion, cinematic glow.", 0], | |
| ["Close-up 3D animated scene of a short, fluffy monster with large, wide eyes and an open mouth, kneeling beside a melting red candle, gazing at the flickering flame in awe. The creature's soft, plush-like fur glows under warm, dramatic lighting, emphasizing its innocent and curious expression. Its playful, hunched posture suggests wonder and discovery, as if experiencing fire for the first time. Behind, a cozy, warmly lit room features a crackling fireplace, soft rugs, and wooden furniture blurred in the background. Rich amber, crimson, and golden hues enhance the magical, intimate atmosphere. Gentle flickers of candlelight animate the scene with subtle realism. Smooth Pixar-style rendering, shallow depth of field, eye-level angle.", 0], | |
| ["A vibrant anime-style illustration in thick, expressive brushwork depicting a young man in his 20s with short messy black hair and warm brown eyes, deeply engrossed in reading a classic leather-bound book. He sits casually on a fluffy white cloud floating in a radiant evening sky, one leg crossed over the other, wearing a simple white t-shirt and blue jeans. The background bursts with soft cotton-like clouds bathed in a golden sunset glow, casting a warm orange hue across the dreamy, ethereal atmosphere. Rendered in rich, painterly textures with a medium shot from a slightly downward angle, emphasizing his serene focus and the tranquil skyward setting.", 0], | |
| ["Classic cinematic movie trailer, a determined 30-year-old space explorer journeys across a vast salt desert under a boundless blue sky. He wears a striking red wool knitted motorcycle helmet that glints in the harsh sunlight, contrasting vividly against the pale, cracked terrain. Shot on 35mm film with rich, saturated colors and fine grain texture, the scene captures sweeping desert vistas, shimmering salt flats, and endless horizons. Dynamic medium shots transition to sweeping overhead angles, emphasizing his resilience and the scale of his solitary adventure. Dramatic lighting and slow-motion details highlight every step forward.", 0], | |
| ], | |
| inputs=[prompt, seed], | |
| fn=generate, | |
| outputs=[streaming_video, seed_out, status_display], | |
| # The handler's final yield sets the video slot to None, and cached | |
| # None outputs break example clicks in current Gradio — run on click. | |
| cache_examples=False, | |
| run_on_click=True, | |
| ) | |
| generate_btn.click( | |
| fn=generate, | |
| inputs=[prompt, seed, num_blocks], | |
| outputs=[streaming_video, seed_out, status_display], | |
| api_name="generate", | |
| ) | |
| demo.queue() | |
| # theme / css moved to launch(): the Gradio 6 Blocks constructor deprecates them | |
| # (they still work there, but emit a UserWarning into the Space log). | |
| demo.launch(theme=gr.themes.Citrus(), css=CSS, mcp_server=True) |