multimodalart's picture
multimodalart HF Staff
Upload app.py with huggingface_hub
6ba71e5 verified
Raw
History Blame Contribute Delete
7 kB
import os
# Set expandable segments for memory pressure (video DiT)
os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")
import spaces # MUST come before torch / diffusers / transformers
import torch
import gradio as gr
import tempfile
import time
import numpy as np
import imageio
from omegaconf import OmegaConf
from einops import rearrange
from pipeline import CausalInferencePipeline
from wan.modules.sparse_attention import calculate_chunk_sparsities
MODEL_ID = "Wan-AI/Wan2.1-T2V-1.3B"
LF_CKPT_ID = "mack-williams/Light-Forcing"
# --- Model loading (module scope, eagerly on cuda) ---
from huggingface_hub import snapshot_download
if not os.path.exists("wan_models/Wan2.1-T2V-1.3B"):
print("Downloading Wan2.1-T2V-1.3B base model...")
snapshot_download(
repo_id=MODEL_ID,
local_dir="wan_models/Wan2.1-T2V-1.3B",
)
lf_ckpt_dir = "checkpoints"
os.makedirs(lf_ckpt_dir, exist_ok=True)
lf_ckpt_path = os.path.join(lf_ckpt_dir, "short_video_gen.pt")
if not os.path.exists(lf_ckpt_path):
print("Downloading Light-Forcing checkpoint...")
from huggingface_hub import hf_hub_download
hf_hub_download(
repo_id=LF_CKPT_ID,
filename="short_video_gen.pt",
local_dir=lf_ckpt_dir,
)
# Load config
config = OmegaConf.load("configs/light_forcing_short.yaml")
default_config = OmegaConf.load("configs/default_config.yaml")
config = OmegaConf.merge(default_config, config)
# Calculate sparse attention sparsity schedule
num_frame_per_block = getattr(config, "num_frame_per_block", 1)
model_kwargs = dict(getattr(config, "model_kwargs", {}) or {})
local_attn_size = model_kwargs.get("local_attn_size", 21)
sparse_config = model_kwargs.get("sparse_config", {}) or {}
NUM_OUTPUT_FRAMES = 21
sparsity_list = calculate_chunk_sparsities(
NUM_OUTPUT_FRAMES, num_frame_per_block, local_attn_size, sparse_config
)
if sparsity_list:
sparse_config["sparsity_list"] = sparsity_list
print(f"Sparsity list: {sparsity_list}")
# Initialize pipeline (CausalInferencePipeline handles all model init internally)
pipeline = CausalInferencePipeline(config, device="cuda")
# Load Light Forcing checkpoint
state_dict = torch.load(lf_ckpt_path, map_location="cpu", weights_only=False)
pipeline.generator.load_state_dict(state_dict["generator_ema"])
pipeline = pipeline.to(dtype=torch.bfloat16)
pipeline.text_encoder.to("cuda")
pipeline.generator.to("cuda")
pipeline.vae.to("cuda")
pipeline.text_encoder.eval()
pipeline.generator.eval()
pipeline.vae.eval()
pipeline.text_encoder.requires_grad_(False)
pipeline.generator.requires_grad_(False)
pipeline.vae.requires_grad_(False)
print("Model loaded successfully!")
@spaces.GPU(duration=60)
def generate(
prompt: str,
seed: int = 42,
num_output_frames: int = 21,
progress=gr.Progress(track_tqdm=True),
):
"""Generate a short video from a text prompt using Light Forcing sparse attention."""
if not prompt.strip():
return None, "Please enter a prompt."
torch.manual_seed(seed)
torch.set_grad_enabled(False)
start_time = time.time()
# Generate noise (bfloat16 to match model)
noise = torch.randn(
[1, num_output_frames, 16, 64, 96],
device="cuda",
dtype=torch.bfloat16,
)
# Run inference using the pipeline's built-in method
video = pipeline.inference(
noise=noise,
text_prompts=[prompt],
return_latents=False,
profile=False,
low_memory=False,
)
# video: [b, t, c, h, w] in [0, 1]
video = rearrange(video, 'b t c h w -> b t h w c').cpu()
# Save as MP4
video_np = (video[0].numpy() * 255).clip(0, 255).astype(np.uint8)
with tempfile.NamedTemporaryFile(suffix=".mp4", delete=False) as f:
output_path = f.name
imageio.mimsave(output_path, video_np, fps=16, quality=8)
elapsed = time.time() - start_time
print(f"Generation completed in {elapsed:.2f}s")
return output_path, f"Generated {num_output_frames} frames in {elapsed:.1f}s"
CSS = """
#col-container { max-width: 900px; margin: 0 auto; }
"""
with gr.Blocks() as demo:
gr.Markdown("""
# Light Forcing: Accelerating Autoregressive Video Diffusion via Sparse Attention
Generate short videos from text prompts using the Light Forcing sparse attention method
on top of the Wan2.1-T2V-1.3B autoregressive video diffusion model.
[Paper](https://arxiv.org/abs/2602.04789) | [GitHub](https://github.com/chengtao-lv/LightForcing) | [Model](https://huggingface.co/mack-williams/Light-Forcing)
""")
with gr.Column(elem_id="col-container"):
with gr.Row():
prompt = gr.Textbox(
show_label=False,
placeholder="Describe the video you want to generate...",
container=False,
scale=4,
)
run_btn = gr.Button("Generate", variant="primary", scale=1)
output_video = gr.Video(label="Generated Video")
status_text = gr.Textbox(label="Status", interactive=False)
with gr.Accordion("Advanced settings", open=False):
seed = gr.Number(label="Seed", value=42, precision=0)
num_frames = gr.Slider(
label="Number of frames",
minimum=3,
maximum=21,
value=21,
step=3,
info="21 frames ≈ 5 seconds at 16fps (short video mode)",
)
gr.Examples(
examples=[
["A stylish woman walks down a Tokyo street filled with warm glowing neon and animated city signage. She wears a black leather jacket, a long red dress, and black boots, and carries a black purse. She wears sunglasses and red lipstick. She walks confidently and casually."],
["Several giant wooly mammoths approach treading through a snowy meadow, their long wooly fur lightly blows in the wind as they walk, snow covered trees and dramatic snow capped mountains in the distance, mid afternoon light with wispy clouds and a sun high in the distance creates a warm glow, the low camera view is stunning capturing the large furry mammal with beautiful photography, depth of field."],
["A movie trailer featuring the adventures of the 30 year old space man wearing a red wool knitted motorcycle helmet, blue sky, salt desert, cinematic style, shot on 35mm film, vivid colors."],
["Drone view of waves crashing against the rugged cliffs along Big Sur's garay point beach. The crashing blue waters create white-tipped waves, while the golden light of the setting sun illuminates the rocky shore."],
],
inputs=[prompt],
outputs=[output_video, status_text],
fn=generate,
cache_examples=True,
cache_mode="lazy",
)
run_btn.click(
fn=generate,
inputs=[prompt, seed, num_frames],
outputs=[output_video, status_text],
)
demo.launch(mcp_server=True, theme=gr.themes.Citrus(), css=CSS)