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)