Spaces:
Running on Zero
Running on Zero
File size: 18,016 Bytes
03b1c55 1ed0eb8 03b1c55 1ed0eb8 03b1c55 1ed0eb8 03b1c55 1ed0eb8 03b1c55 1ed0eb8 03b1c55 1ed0eb8 03b1c55 1ed0eb8 03b1c55 1ed0eb8 03b1c55 1ed0eb8 03b1c55 1ed0eb8 03b1c55 1ed0eb8 03b1c55 1ed0eb8 03b1c55 1ed0eb8 03b1c55 1ed0eb8 6605f04 03b1c55 1ed0eb8 03b1c55 1ed0eb8 03b1c55 1ed0eb8 03b1c55 1ed0eb8 03b1c55 1ed0eb8 30b6bc4 03b1c55 1ed0eb8 03b1c55 1ed0eb8 03b1c55 6d97beb 03b1c55 6d97beb 1ed0eb8 03b1c55 1ed0eb8 03b1c55 1ed0eb8 03b1c55 1ed0eb8 03b1c55 1ed0eb8 03b1c55 1ed0eb8 6d97beb | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383 384 385 386 387 388 389 390 | 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) |