akhaliq's picture
akhaliq HF Staff
Return serveable file dicts from fn nodes; enable hf_oauth
6adf256
Raw History Blame Contribute Delete
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:
@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)