"""Ahead-of-time (torch.export + AOTInductor) compiler for Zing-0.5. Companion Space to https://huggingface.co/spaces/hugging-apps/zing-0-5-world-model. The demo Space spends all of its GPU time in 30 identical `WanAttentionBlock`s (5 forward passes per generated block: 4 DMD denoising steps + the KV-cache commit). This Space compiles that block once, ahead of time, and publishes the resulting `.pt2` archive to a model repo laid out the way `spaces.aoti_blocks_load` expects: /WanAttentionBlock/package.pt2 The demo then calls spaces.aoti_blocks_load(PIPELINE.generator, "") at start-up and every `WanAttentionBlock` instance runs the precompiled graph instead of eager PyTorch -- no compilation at runtime, no warm-up on the first rollout. Convention notes (this follows the `spaces` package, not a bespoke scheme): * `spaces.aoti_capture(block)` records the exact args the block is called with during a real rollout, so the exported graph is traced against real tensors. * `spaces.aoti_compile_and_save(dir, exported)` writes `dir/root/package.pt2`; that file is uploaded as `WanAttentionBlock/package.pt2`, which is the `{block class name}/package.pt2` layout `spaces.aoti_blocks_load` downloads (cf. `zerogpu-aoti/FLUX.1`, `zerogpu-aoti/Wan2`). * Weights are *not* baked into the archive; `spaces.aoti_patch` feeds each block's own `state_dict()` to the compiled model at call time, so one artifact serves all 30 layers. """ import os os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True") os.environ.setdefault("PYTORCH_ALLOC_CONF", "expandable_segments:True") import spaces # noqa: E402 — must precede torch / any CUDA-touching import import traceback # noqa: E402 from pathlib import Path # noqa: E402 import gradio as gr # noqa: E402 import torch # noqa: E402 from huggingface_hub import HfApi, snapshot_download # noqa: E402 from zing_v0_5.config import load_config, with_cache_window # noqa: E402 from zing_v0_5.pipeline import InferencePipeline # noqa: E402 from zing_v0_5.processor import MessageProcessor # noqa: E402 MODEL_ID = "seedleap/zing-0.5" ROOT = Path(__file__).parent # Must match the demo Space exactly — the exported graph is traced against the # shapes this sliding-window configuration produces. LOCAL_ATTN_SIZE = 33 SINK_SIZE = 5 BLOCK_CLASS = "WanAttentionBlock" AOTI_REPO = os.environ.get("ZING_AOTI_REPO", "hugging-apps/zing-0-5-world-model-aoti") DEMO_SPACE = "hugging-apps/zing-0-5-world-model" RESOLUTIONS = { "832 × 480 (demo default)": (832, 480), "1088 × 608": (1088, 608), "1248 × 704 (native)": (1248, 704), } DEFAULT_RESOLUTION = "832 × 480 (demo default)" # 45 pixel frames == 12 latent frames == 3 blocks. Block 0 runs with a cold KV # cache (and is deliberately excluded from the AoTI path in the demo), so the # capture lands on block 1, i.e. a *warm* cache — exactly the shape family the # compiled graph has to serve. CAPTURE_FRAMES = 45 torch.set_grad_enabled(False) torch.backends.cuda.matmul.allow_tf32 = True torch.backends.cudnn.allow_tf32 = True print("Downloading Zing-0.5 weights …", flush=True) SNAPSHOT = Path(snapshot_download(MODEL_ID)) CHECKPOINT = SNAPSHOT / "generator" / "model.pt" PRETRAINED = SNAPSHOT / "pretrained" CONFIG = with_cache_window(load_config(ROOT / "config" / "zing.yaml"), LOCAL_ATTN_SIZE, SINK_SIZE) print("Building the Zing-0.5 pipeline …", flush=True) PIPELINE = InferencePipeline(CONFIG, PRETRAINED, CHECKPOINT) PROCESSOR = MessageProcessor(CONFIG, PIPELINE.encode_reference) print("Pipeline ready.", flush=True) def _request(prompt: str, resolution: str): width, height = RESOLUTIONS.get(resolution, RESOLUTIONS[DEFAULT_RESOLUTION]) return PROCESSOR.process( { "sample_id": "zing_aoti_capture", "messages": [ {"role": "user", "type": "text", "content": prompt}, { "role": "target", "type": "video", "reference_frame_count": 0, "output": {"frames": CAPTURE_FRAMES, "height": height, "width": width}, "controls": [], }, ], } ) def _describe(args) -> str: names = ( "hidden", "time_embedding", "rope", "history_key", "history_value", "cross_key", "cross_value", "context_lengths", ) return "\n".join( f" {name:<16} {tuple(value.shape)} {value.dtype}" for name, value in zip(names, args) if isinstance(value, torch.Tensor) ) @spaces.GPU(duration=1500) def _compile(prompt: str, resolution: str, max_autotune: bool) -> tuple[bytes, str]: """Capture → `torch.export` → AOTInductor. Returns the `.pt2` bytes.""" log: list[str] = [] request = _request(prompt, resolution) block = PIPELINE.generator.blocks[0] # The demo routes cold-cache (block 0) calls through the *unbound* forward, # so this capture fires on the first warm-cache call — a non-empty history. with spaces.aoti_capture(block) as call: for _ in PIPELINE.stream(request): pass if not call.args: raise RuntimeError("no WanAttentionBlock call was captured") args = tuple(call.args) log.append(f"captured call ({len(args)} tensors):\n{_describe(args)}") hidden, _, _, history_key, _, cross_key, *_ = args # Three axes move at runtime: the token count of a block (resolution), the # length of the KV history (grows until the sliding window caps it) and the # prompt length (changes on every mid-session rewrite). sequence = torch.export.Dim("sequence", min=256, max=8192) history = torch.export.Dim("history", min=8, max=32768) context = torch.export.Dim("context", min=8, max=4096) dynamic_shapes = ( {1: sequence}, # hidden [1, seq, dim] {1: sequence}, # time_embedding [1, seq, 6, dim] {0: sequence}, # rope [seq, dim/2, 2] {1: history}, # history_key [1, hist, heads, head_dim] {1: history}, # history_value {0: context}, # cross_key [ctx, heads, head_dim] {0: context}, # cross_value {}, # context_lengths [1] ) log.append( f"dynamic axes: sequence={hidden.shape[1]}, history={history_key.shape[1]}, " f"context={cross_key.shape[0]}" ) print("Exporting WanAttentionBlock …", flush=True) exported = torch.export.export( block, args, {}, dynamic_shapes=dynamic_shapes, strict=False ) log.append("torch.export: ok") inductor_configs = {"max_autotune": True} if max_autotune else {} print("AOTInductor compile …", flush=True) package_dir = Path("/tmp/zing-aoti-package") spaces.aoti_compile_and_save(package_dir, exported, inductor_configs) package = package_dir / "root" / "package.pt2" payload = package.read_bytes() log.append(f"AOTInductor: ok — package.pt2 is {len(payload) / 1e6:.1f} MB") return payload, "\n".join(log) def compile_and_publish(prompt: str, resolution: str, max_autotune: bool, publish: bool): try: payload, log = _compile(prompt, resolution, max_autotune) except Exception: # pragma: no cover - surfaced in the UI return "❌ compilation failed\n\n" + traceback.format_exc() lines = [log] if not publish: return "✅ compiled (not published — tick “publish” to upload)\n\n" + "\n".join(lines) token = os.environ.get("HF_TOKEN") if not token: return "⚠️ compiled, but HF_TOKEN is not set on this Space.\n\n" + "\n".join(lines) local = Path("/tmp/zing-aoti-upload") / BLOCK_CLASS local.mkdir(parents=True, exist_ok=True) (local / "package.pt2").write_bytes(payload) api = HfApi(token=token) api.create_repo(AOTI_REPO, repo_type="model", exist_ok=True) api.upload_folder( repo_id=AOTI_REPO, repo_type="model", folder_path=str(local.parent), commit_message=f"AOTI {BLOCK_CLASS} ({resolution}, max_autotune={max_autotune})", ) lines.append( f"published → https://huggingface.co/{AOTI_REPO}/blob/main/{BLOCK_CLASS}/package.pt2" ) return "✅ compiled and published\n\n" + "\n".join(lines) with gr.Blocks(title="Zing-0.5 · AoTI compiler") as demo: gr.Markdown( f""" # Zing-0.5 — ahead-of-time compiler Compiles `{BLOCK_CLASS}` (the 30× repeated transformer block of [`{MODEL_ID}`](https://huggingface.co/{MODEL_ID})) with **`torch.export` + AOTInductor** and publishes the archive to [`{AOTI_REPO}`](https://huggingface.co/{AOTI_REPO}) as `{BLOCK_CLASS}/package.pt2`. [`{DEMO_SPACE}`](https://huggingface.co/spaces/{DEMO_SPACE}) loads it at start-up with `spaces.aoti_blocks_load(...)`, so the live demo never compiles anything at runtime. """ ) with gr.Row(): with gr.Column(): prompt = gr.Textbox( label="Capture prompt", value="In first-person perspective, walking down a sunlit forest path", lines=2, ) resolution = gr.Dropdown( list(RESOLUTIONS), value=DEFAULT_RESOLUTION, label="Capture resolution" ) max_autotune = gr.Checkbox( value=False, label="max-autotune", info="Slower to compile; benchmarks kernels for the captured shape.", ) publish = gr.Checkbox(value=True, label=f"publish to {AOTI_REPO}") run = gr.Button("Compile & publish", variant="primary") with gr.Column(): report = gr.Textbox(label="Report", lines=22) run.click( compile_and_publish, inputs=[prompt, resolution, max_autotune, publish], outputs=report, api_name="compile", ) demo.queue().launch(theme=gr.themes.Citrus())