"""WorldCrafter-Fast on ZeroGPU with per-action streaming state. The checkpoint loader retains the Space's disk and CUDA adaptations. Each button request uses the shared Fast sampler and saves CPU state for the next GPU allocation; models remain shared while user histories stay separate. """ 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 = 384, 640 # -------------------------------------------------------------------------------------- # 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: if os.environ.get("WORLDCRAFTER_UI_ONLY") != "1": 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) # One action is one GPU job. Session tensors remain on CPU between jobs. @spaces.GPU(duration=65, size="xlarge") def generate_step(image, prompt, seed, event, radius, session, progress_id=None): if MODEL is None: raise gr.Error("The video model is unavailable. Please check the runtime logs.") from interactive import step return step(MODEL, image, prompt, seed, event, radius, session, progress_id) from space_ui import build_demo demo, CSS = build_demo(generate_step, LOAD_ERROR) if __name__ == "__main__": demo.launch(server_name="0.0.0.0", share=os.environ.get("WORLDCRAFTER_SHARE") == "1", theme=gr.themes.Base(primary_hue="lime"), css=CSS)