multimodalart's picture
multimodalart HF Staff
Upload app.py with huggingface_hub
6605f04 verified
Raw History Blame Contribute Delete
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
@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"<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)