"""WorldCrafter-Fast on ZeroGPU. Mirrors the official inference path (`python inference.py --model-type fast`) from https://github.com/TencentARC/WorldCrafter: the vendored `worldcrafter` package is the authors' own code, and `WorldCrafter.generate(...)` is called with the same arguments the CLI uses. Two deviations, both forced by the ZeroGPU runtime, are documented at `_SeparateBranches` and `_no_set_device` below. """ import os os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True") # After packing, ZeroGPU auto-prunes any Hub blob still mmap-backing a packed tensor and # `lstat()`s it to report the reclaimed size. `_release()` below already deletes those # blobs - much earlier, which is the only way this checkpoint fits the storage quota at # all - so the post-pack `lstat` would hit a ` (deleted)` path and abort startup. # Pruning is ours to do here, so switch the built-in pass off by pointing it at nothing. os.environ.setdefault("ZEROGPU_MMAP_AUTOPRUNE_PATTERN", "/zerogpu-autoprune-disabled/*") import spaces # noqa: E402 (must precede torch / CUDA-touching imports) import torch # noqa: E402 import gradio as gr # noqa: E402 import gc # noqa: E402 import json # noqa: E402 import subprocess # noqa: E402 import tempfile # noqa: E402 import time # noqa: E402 import traceback # noqa: E402 from pathlib import Path # noqa: E402 from huggingface_hub import hf_hub_download, snapshot_download # noqa: E402 MODEL_ID = "TencentARC/WorldCrafter-Fast" HERE = Path(__file__).parent EXAMPLES_DIR = HERE / "examples" HEIGHT, WIDTH, FPS = 384, 640, 16 CHUNK_FRAMES = 33 MAX_CHUNKS = 6 DEFAULT_CHUNKS = 3 MAX_SEED = 2**31 - 1 NEGATIVE_PROMPT = (EXAMPLES_DIR / "negative_prompt.txt").read_text(encoding="utf-8").strip() # -------------------------------------------------------------------------------------- # ZeroGPU deviations from the reference loader # -------------------------------------------------------------------------------------- class _SeparateBranches: """Stand-in for `worldcrafter.fast.resident.ResidentBranches` on ZeroGPU. `ResidentBranches` is a pure *memory* optimisation, not a modelling component: it keeps ONE materialised bf16 tensor per shared parameter plus reversible packed integer bit-pattern deltas, and flips between the two distilled experts by launching triton kernels captured into CUDA graphs. Building it therefore needs a live CUDA context, triton JIT and `torch.cuda.CUDAGraph()` *at load time* - none of which exist in a ZeroGPU Space's main process, where models are loaded before any GPU is attached. Here both experts stay fully materialised in bf16, so there is nothing to switch and this object only keeps the branch bookkeeping that `worldcrafter.fast.sampling` and the routing hooks in `model_loading.load_fast` read (`switch`, `active`, `switches`). `sampling.sample_fast` selects the module itself via `stage_transformers[2 if branch == "old" else 0]`, so the weights each step sees are bit-identical to the reference; the only cost is the ~28.6 GB of VRAM that sharing would have saved (hence `xlarge`). """ def __init__(self, early, late, use_graph=True): early_params = dict(early.named_parameters()) late_params = dict(late.named_parameters()) if set(early_params) != set(late_params): raise ValueError("Invalid branch weight layout") shared_bytes = 0 independent = 0 for name, left in early_params.items(): right = late_params[name] if left.dtype != right.dtype or left.shape != right.shape: raise ValueError(name) if ".lora_" in name or ".cam_self_attn." in name or left.dtype != torch.bfloat16: independent += 1 else: shared_bytes += left.numel() * left.element_size() self.active = "equal" self.switches = 0 self.report = dict( implementation="separate_materialized_branches (ZeroGPU)", reason="ResidentBranches needs triton + CUDA graphs at load time", materialized_bf16_bytes_per_branch=shared_bytes, independent_parameter_tensors=independent, lossless_bf16_bit_patterns=True, gpu_only_switch=False, cuda_graph=False, ) def switch(self, branch): if branch == self.active: return if branch not in ("equal", "old"): raise ValueError(branch) self.active = branch self.switches += 1 def _no_set_device(device=None): """`load_fast` calls `torch.cuda.set_device(...)`; ZeroGPU re-assigns device ids per request, so pinning one at import time is both meaningless and a CUDA-init hazard.""" return None # -------------------------------------------------------------------------------------- # Quota-aware weight staging # # `TencentARC/WorldCrafter-Fast` ships fp32: the two 14.3B experts are 53.2 GiB each and # the UMT5-XXL text encoder is 21.2 GiB, 137 GiB in total. A Space's ephemeral storage is # capped at 150 GB, and ZeroGPU additionally writes every packed bf16 weight back to that # same disk at the startup pack step (~75 GiB here), so a plain `snapshot_download` gets # the workload evicted with "storage limit exceeded". # # So the three oversized components are fetched only for as long as they are being read. # `load_fast` loads them one at a time (`transformer("high")` -> text encoder -> ... -> # `transformer("low")`), so wrapping each class's `from_pretrained` with fetch/release # keeps peak staging at one component (~53 GiB) instead of all of them (137 GiB), and the # weights themselves still travel through the authors' own loading code with the authors' # own fp32 -> bf16 cast, `_keep_in_fp32_modules` included. Nothing about the numerics # changes; only *when* the bytes are on disk does. # # Releasing is safe precisely because every tensor is dtype-cast on load, so no parameter # can still be backed by the checkpoint's mmap when the file is unlinked. # -------------------------------------------------------------------------------------- LAZY_PREFIXES = ("transformer_high_noise/", "transformer_low_noise/", "text_encoder/") _LAZY_FILES: dict[str, int] = {} _SNAPSHOT = Path(".") _STAGING = Path(".") def _disk(label): out = [] for path in (Path.home() / ".cache/huggingface", _STAGING): if path.exists(): try: used = subprocess.run(["du", "-sBG", str(path)], capture_output=True, text=True, timeout=120).stdout.split()[0] except Exception: # noqa: BLE001 used = "?" out.append(f"{path.name}={used}") print(f"[space] disk after {label}: {' '.join(out)}", flush=True) def _fetch(prefix): paths = sorted(p for p in _LAZY_FILES if p.startswith(prefix)) started = time.perf_counter() for rel in paths: blob = Path(hf_hub_download(MODEL_ID, rel)) actual = blob.stat().st_size if actual != _LAZY_FILES[rel]: raise ValueError(f"Incomplete fast checkpoint: {rel} ({actual} bytes)") link = _STAGING / rel link.parent.mkdir(parents=True, exist_ok=True) if not link.exists(): link.symlink_to(blob.resolve()) print(f"[space] fetched {prefix} ({len(paths)} files) in " f"{time.perf_counter() - started:.0f}s", flush=True) def _release(prefix): for rel in sorted(p for p in _LAZY_FILES if p.startswith(prefix)): for link in (_STAGING / rel, _SNAPSHOT / rel): target = link.resolve() if link.is_symlink() else None if link.is_symlink() or link.exists(): link.unlink() if target is not None and target.is_file(): target.unlink() gc.collect() _disk(f"release {prefix}") def _stage_checkpoint(): """Download everything small, and mirror it into a staging root whose manifest no longer claims the three lazily-fetched components are already on disk.""" global _LAZY_FILES, _SNAPSHOT, _STAGING print(f"[space] downloading {MODEL_ID} (small components) ...", flush=True) started = time.perf_counter() _SNAPSHOT = Path( snapshot_download( MODEL_ID, ignore_patterns=[f"{p}*.safetensors" for p in LAZY_PREFIXES], max_workers=16, ) ) print(f"[space] snapshot ready in {time.perf_counter() - started:.0f}s", flush=True) manifest = json.loads((_SNAPSHOT / "manifest.json").read_text()) _LAZY_FILES = { row["path"]: row["bytes"] for row in manifest["files"] if row["path"].startswith(LAZY_PREFIXES) and row["path"].endswith(".safetensors") } if len(_LAZY_FILES) != 17: # 6 + 6 transformer shards, 5 text-encoder shards raise RuntimeError(f"Unexpected fast checkpoint layout: {sorted(_LAZY_FILES)}") _STAGING = Path(tempfile.mkdtemp(prefix="worldcrafter-weights-")) / "WorldCrafter-Fast" _STAGING.mkdir(parents=True) for src in sorted(_SNAPSHOT.rglob("*")): if src.is_dir(): continue dst = _STAGING / src.relative_to(_SNAPSHOT) dst.parent.mkdir(parents=True, exist_ok=True) dst.symlink_to(src.resolve()) # The lazily-fetched shards are byte-size checked in `_fetch` exactly as `load_fast` # would; drop only their rows so the rest of the manifest is still enforced. (_STAGING / "manifest.json").unlink() (_STAGING / "manifest.json").write_text( json.dumps( {**manifest, "files": [r for r in manifest["files"] if r["path"] not in _LAZY_FILES]}, indent=2, ) ) _disk("staging") return _STAGING def _lazy_loader(real_cls, prefix_of): """`from_pretrained` that fetches its component, loads it, then frees the bytes.""" class _Lazy: @staticmethod def from_pretrained(path, *args, **kwargs): prefix = prefix_of(Path(path)) _fetch(prefix) try: started = time.perf_counter() model = real_cls.from_pretrained(path, *args, **kwargs) print(f"[space] loaded {prefix} in {time.perf_counter() - started:.0f}s", flush=True) finally: _release(prefix) return model return _Lazy def _load_model(): import worldcrafter.model_loading as model_loading from worldcrafter import WorldCrafter model_loading.ResidentBranches = _SeparateBranches torch.cuda.set_device = _no_set_device model_loading.WorldCrafterTransformer3DModel = _lazy_loader( model_loading.WorldCrafterTransformer3DModel, lambda p: f"{p.name}/" ) model_loading.UMT5EncoderModel = _lazy_loader( model_loading.UMT5EncoderModel, lambda p: "text_encoder/" ) root = _stage_checkpoint() started = time.perf_counter() model = WorldCrafter.from_pretrained( root, model_type="fast", device="cuda", height=HEIGHT, width=WIDTH, attention_backend="native", enable_compile=False, ) print(f"[space] model assembled in {time.perf_counter() - started:.0f}s", flush=True) print("[space] fast_report: " + json.dumps(model.fast_report, indent=2, default=str), flush=True) _disk("load") return model MODEL = None LOAD_ERROR = None try: MODEL = _load_model() except Exception as exc: # noqa: BLE001 - keep the Space up so logs stay reachable LOAD_ERROR = f"{exc!r}\n{traceback.format_exc()}" print(f"[space] MODEL LOAD FAILED:\n{LOAD_ERROR}", flush=True) # -------------------------------------------------------------------------------------- # Inference # -------------------------------------------------------------------------------------- def _build_camera(actions: str, num_chunks: int, workdir: Path): from worldcrafter.camera import build_trajectory, count_chunks, parse_trajectory try: events, options = parse_trajectory(actions or "") except ValueError as exc: raise gr.Error(f"Invalid camera actions: {exc}") from exc if not events: raise gr.Error("Add at least one camera action, for example `forward1x2`.") available = count_chunks(events) chunks = max(1, min(int(num_chunks), MAX_CHUNKS, available)) try: camera, records = build_trajectory(events, **options) except (ValueError, KeyError) as exc: raise gr.Error(f"Invalid camera actions: {exc}") from exc from worldcrafter.camera import save_trajectory camera_path = save_trajectory( workdir, camera, records, fps=FPS, events=events, options=options ) return camera_path, chunks, available def _run(mode, image_path, prompt, actions, negative_prompt, num_chunks, seed): if MODEL is None: raise gr.Error(f"The model failed to load at startup:\n{LOAD_ERROR}") if not (prompt or "").strip(): raise gr.Error("A prompt is required.") if mode == "i2v" and not image_path: raise gr.Error("Upload a start image, or switch to the Text → video tab.") workdir = Path(tempfile.mkdtemp(prefix="worldcrafter-")) camera_path, chunks, available = _build_camera(actions, num_chunks, workdir) output_path = workdir / "worldcrafter.mp4" started = time.perf_counter() result = MODEL.generate( mode=mode, camera_path=camera_path, output_path=output_path, prompt=prompt.strip(), negative_prompt=(negative_prompt or "").strip(), image_path=Path(image_path) if mode == "i2v" else None, num_chunks=chunks, seed=int(seed), fps=FPS, ) elapsed = time.perf_counter() - started summary = result.summary info = ( f"**{chunks} chunk(s)** · {summary['num_frames']} frames · {WIDTH}x{HEIGHT} @ {FPS} fps " f"· seed {summary['seed']} · {summary['num_inference_steps']} steps, CFG " f"{summary['guidance_scale']:g} · routing {summary.get('first_chunk_routing', '?')} " f"then {summary.get('subsequent_chunk_routing', '?')} · **{elapsed:.1f}s** " f"({elapsed / chunks:.1f}s/chunk)" ) if available > chunks: info += ( f"\n\nThe action script describes {available} chunks; only the first {chunks} " "were rendered. Raise *Chunks to generate* to go further." ) print( f"[space] {mode} done in {elapsed:.1f}s for {chunks} chunk(s) · " f"{torch.cuda.get_device_name()} · peak VRAM " f"{torch.cuda.max_memory_allocated() / 2**30:.1f}/" f"{torch.cuda.get_device_properties(0).total_memory / 2**30:.1f} GiB", flush=True, ) return str(result.video_path), info def _duration(num_chunks, first_chunk_seconds): """Measured on this Space: ~22s to stream the packed weights into VRAM, then 13.4s for a 6-step I2V first chunk (23s for the 12-step T2V one) and 11.4s per chunk after it. Kept tight on purpose - `duration` is charged against each visitor's quota.""" chunks = max(1, min(int(num_chunks), MAX_CHUNKS)) return int(round(1.15 * (22.0 + first_chunk_seconds + (chunks - 1) * 11.4))) def _duration_i2v(image, prompt, actions, negative_prompt=None, num_chunks=DEFAULT_CHUNKS, seed=42, *args, **kwargs): return _duration(num_chunks, 13.5) def _duration_t2v(prompt, actions, negative_prompt=None, num_chunks=DEFAULT_CHUNKS, seed=42, *args, **kwargs): return _duration(num_chunks, 23.0) @spaces.GPU(duration=_duration_i2v, size="xlarge") def generate_i2v( image: str, prompt: str, actions: str, negative_prompt: str = NEGATIVE_PROMPT, num_chunks: int = DEFAULT_CHUNKS, seed: int = 42, progress=gr.Progress(track_tqdm=True), ): """Explore the scene in an image with a scripted camera, as a video. Args: image: Start frame; it is resized to 640x384 and becomes the video's first frame. prompt: Description of the scene and of what the camera should find in it. actions: Camera action script, one action per 33-frame chunk (e.g. `forward1x2`). negative_prompt: Attributes to suppress. num_chunks: How many 33-frame chunks of the action script to render. seed: Random seed. Returns: The generated mp4 path and a markdown run summary. """ return _run("i2v", image, prompt, actions, negative_prompt, num_chunks, seed) @spaces.GPU(duration=_duration_t2v, size="xlarge") def generate_t2v( prompt: str, actions: str, negative_prompt: str = NEGATIVE_PROMPT, num_chunks: int = DEFAULT_CHUNKS, seed: int = 42, progress=gr.Progress(track_tqdm=True), ): """Generate a world from text alone and explore it with a scripted camera. Args: prompt: Description of the scene to create and explore. actions: Camera action script, one action per 33-frame chunk (e.g. `forward1x2`). negative_prompt: Attributes to suppress. num_chunks: How many 33-frame chunks of the action script to render. seed: Random seed. Returns: The generated mp4 path and a markdown run summary. """ return _run("t2v", None, prompt, actions, negative_prompt, num_chunks, seed) # -------------------------------------------------------------------------------------- # Workflow entry points (bound as callable nodes on the gr.Workflow canvas) # -------------------------------------------------------------------------------------- def _as_path(value): """Canvas reference/operator values for media ports can be `{path, url}` dicts.""" if isinstance(value, dict): return value.get("path") or value.get("url") return value def _as_file(value): """`call_fn` JSON-serializes bound-function results verbatim (only `call_space` rewrites local paths into serveable file dicts), so media outputs must be returned in the `{path, url, is_file}` shape the canvas renders. The video lives under the system tempdir, which `Workflow.launch()` adds to `allowed_paths`.""" if not isinstance(value, str) or not os.path.exists(value): return value try: from gradio_client import utils as client_utils encoded = client_utils.encode_file_path(value) except (ImportError, AttributeError): import urllib.parse encoded = urllib.parse.quote(os.path.abspath(value)) return {"path": value, "url": "/gradio_api/file=" + encoded, "is_file": True} def i2v(image, prompt, actions, negative_prompt=NEGATIVE_PROMPT, num_chunks=DEFAULT_CHUNKS, seed=42): """Image → video node: explore an uploaded scene with a scripted camera.""" video, info = generate_i2v( _as_path(image), prompt, actions, negative_prompt or NEGATIVE_PROMPT, int(num_chunks), int(seed), ) return _as_file(video), info def t2v(prompt, actions, negative_prompt=NEGATIVE_PROMPT, num_chunks=DEFAULT_CHUNKS, seed=42): """Text → video node: generate a world from text and explore it.""" video, info = generate_t2v( prompt, actions, negative_prompt or NEGATIVE_PROMPT, int(num_chunks), int(seed), ) return _as_file(video), info CSS = """ #col-container { max-width: 1180px; margin: 0 auto; } .dark .gradio-container { color: var(--body-text-color); } """ # The visual pipeline lives in `workflow.json` next to this script: two disconnected # pipelines (Image → video and Text → video), each wiring prompt/actions/settings # references into the bound `i2v` / `t2v` function nodes and out to video + summary # subjects. Edit it on the canvas (write-access URL) or by hand; `bind=` keys must # match the operator nodes' `"fn"` values. demo = gr.Workflow( graph="workflow.json", bind={"i2v": i2v, "t2v": t2v}, ) if __name__ == "__main__": demo.launch(theme=gr.themes.Citrus(), css=CSS, mcp_server=True)