Spaces:
Running on Zero
Running on Zero
Download app.py from akhaliq/worldcrafter-demo-workflow: direct link, hf CLI and curl.
- Browser
- Download file 20.1 kB
-
https://huggingface.co/spaces/akhaliq/worldcrafter-demo-workflow/resolve/main/app.py
- Command line
-
hf download hf://spaces/akhaliq/worldcrafter-demo-workflow/app.py
-
curl -L -o app.py https://huggingface.co/spaces/akhaliq/worldcrafter-demo-workflow/resolve/main/app.py
20.1 kB
| """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 `<blob> (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: | |
| 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) | |
| 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) | |
| 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) | |