multimodalart's picture
multimodalart HF Staff
diagnostics bootstrap: vendored worldcrafter package, deps, examples
5cfafaa verified
Raw History Blame Contribute Delete
23.6 kB
from __future__ import annotations
import hashlib
import json
from dataclasses import dataclass
from collections.abc import Callable
from pathlib import Path
import numpy as np
import torch
from diffusers.models import AutoencoderKLWan
from diffusers.utils import export_to_video, load_image
from transformers import AutoTokenizer, UMT5EncoderModel
from .output import sha256, save_chunk_state, assemble_resumed_video
from .diffusers import (
WorldCrafterPipeline,
WorldCrafterScheduler,
WorldCrafterTransformer3DModel,
)
from .kernels import (
replace_all_norms_with_flash_norms,
replace_rmsnorm_with_fp32,
replace_rope_with_flash_rope,
)
from .repencoder import (
RepEncoder,
RepEncoderInferenceMemoryProvider,
RepEncoderInferenceProviderConfig,
)
from .ucpe.bridge import (
enable_ucpe_inference_sdpa_attention,
load_ucpe_camera_adapter_weights,
patch_worldcrafter_transformer_ucpe,
)
CAMERA_CHUNK_FRAMES = 33
MODEL_HEIGHT = 384
MODEL_WIDTH = 640
@dataclass(frozen=True)
class InferenceResult:
video_path: Path
summary_path: Path
summary: dict[str, object]
def load_camera(path: Path, num_chunks: int | None = None) -> np.ndarray:
pose = np.load(path, allow_pickle=False)
if pose.ndim == 3:
pose = pose[None]
if pose.ndim != 4 or pose.shape[-2:] not in ((3, 4), (4, 4)):
raise ValueError(f"camera must be [B,T,3,4] or [B,T,4,4], got {pose.shape}")
if pose.shape[0] != 1 or pose.shape[1] % CAMERA_CHUNK_FRAMES:
raise ValueError(
"WorldCrafter requires one camera trajectory containing complete 33-frame chunks"
)
if not np.issubdtype(pose.dtype, np.floating) or not np.isfinite(pose).all():
raise ValueError("camera must contain finite floating-point c2w matrices")
rotation = np.asarray(pose[..., :3, :3], dtype=np.float64)
gram = np.swapaxes(rotation, -1, -2) @ rotation
if not np.allclose(gram, np.eye(3), atol=5e-3, rtol=0.0):
raise ValueError("camera rotations are not orthonormal")
if not np.allclose(np.linalg.det(rotation), 1.0, atol=5e-3, rtol=0.0):
raise ValueError("camera rotations must have determinant +1")
if num_chunks is not None:
frames = int(num_chunks) * CAMERA_CHUNK_FRAMES
if pose.shape[1] < frames:
raise ValueError(
f"camera has {pose.shape[1] // CAMERA_CHUNK_FRAMES} chunks, "
f"but {num_chunks} were requested"
)
pose = pose[:, :frames]
return np.ascontiguousarray(pose)
def validate_weights(model_path: Path) -> dict[str, Path]:
root = model_path.expanduser().resolve()
config_path = root / "inference_config.json"
config = json.loads(config_path.read_text()) if config_path.is_file() else {}
shared = (root / config.get("shared_components", ".")).resolve()
required = {
"root": root,
"transformer": root / "transformer",
"adapter": root / "adapter",
"repencoder": shared / "repencoder",
"vae": shared / "vae",
"scheduler": shared / "scheduler",
"text_encoder": shared / "text_encoder",
"tokenizer": shared / "tokenizer",
}
missing = [str(path) for path in required.values() if not path.exists()]
for filename in (
required["adapter"] / "camera_adapter.pth",
required["adapter"] / "pytorch_lora_weights.safetensors",
required["repencoder"] / "model.safetensors",
required["repencoder"] / "config.json",
required["repencoder"] / "manifest.json",
):
if not filename.is_file():
missing.append(str(filename))
if missing:
raise FileNotFoundError(
"WorldCrafter-Base is incomplete: " + ", ".join(missing)
)
return required
def configure_attention(
transformer: WorldCrafterTransformer3DModel, backend: str
) -> str:
if backend != "auto":
transformer.set_attention_backend(backend)
return backend
for candidate in ("native", "_flash_3_hub", "flash_hub"):
try:
transformer.set_attention_backend(candidate)
return candidate
except (ImportError, RuntimeError, ValueError):
continue
raise RuntimeError("no supported attention backend is available")
def load_model_adapter(
pipe: WorldCrafterPipeline, adapter_path: Path
) -> dict[str, object]:
from diffusers.loaders.peft import _SET_ADAPTER_SCALE_FN_MAPPING
_SET_ADAPTER_SCALE_FN_MAPPING.setdefault(
WorldCrafterTransformer3DModel.__name__, lambda _model_class, weights: weights
)
state = WorldCrafterPipeline.lora_state_dict(str(adapter_path))
transformer_keys = [key for key in state if key.startswith("transformer.")]
if not transformer_keys:
raise RuntimeError("adapter does not contain transformer low-rank weights")
name = "worldcrafter"
pipe.load_lora_weights(str(adapter_path), adapter_name=name)
pipe.set_adapters([name], adapter_weights=[1.0])
return {"name": name, "tensor_keys": len(transformer_keys)}
class WorldCrafter:
def __init__(
self,
*,
pipeline: WorldCrafterPipeline,
memory_provider: RepEncoderInferenceMemoryProvider,
model_path: Path,
device: torch.device,
attention_backend: str,
adapter_load: dict[str, object],
height: int,
width: int,
) -> None:
self.model_type = "base"
self.pipeline = pipeline
self.memory_provider = memory_provider
self.model_path = model_path
self.device = device
self.attention_backend = attention_backend
self.adapter_load = adapter_load
self.height = height
self.width = width
@classmethod
def from_pretrained(
cls,
model_path: Path,
*,
model_type: str = "base",
device: str = "cuda:0",
height: int = MODEL_HEIGHT,
width: int = MODEL_WIDTH,
seed: int = 42,
memory_fov_h_deg: float = 100.0,
memory_fov_v_deg: float = 71.13349068444832,
memory_fov_samples_per_axis: int = 10,
attention_backend: str = "native",
enable_compile: bool = False,
) -> "WorldCrafter":
if model_type == "fast":
from .model_loading import load_fast
return load_fast(
cls,
model_path,
device=device,
height=height,
width=width,
seed=seed,
memory_fov_h_deg=memory_fov_h_deg,
memory_fov_v_deg=memory_fov_v_deg,
memory_fov_samples_per_axis=memory_fov_samples_per_axis,
attention_backend=attention_backend,
enable_compile=enable_compile,
)
if model_type != "base":
raise ValueError(f"Unknown model type: {model_type}")
torch_device = torch.device(device)
if torch_device.type != "cuda" or not torch.cuda.is_available():
raise RuntimeError("WorldCrafter inference requires CUDA")
if (height, width) != (MODEL_HEIGHT, MODEL_WIDTH):
raise ValueError("WorldCrafter-Base is fixed to 384x640 inference")
torch.cuda.set_device(torch_device)
paths = validate_weights(model_path)
enable_ucpe_inference_sdpa_attention()
repencoder = RepEncoder.from_pretrained(
paths["repencoder"],
device=torch_device,
compute_dtype="bf16",
target_microbatch=4,
)
memory_provider = RepEncoderInferenceMemoryProvider(
repencoder,
RepEncoderInferenceProviderConfig(
seed=seed,
trajectory_fov_horizontal_fov_degrees=memory_fov_h_deg,
trajectory_fov_vertical_fov_degrees=memory_fov_v_deg,
trajectory_fov_samples_per_axis=memory_fov_samples_per_axis,
),
)
transformer = WorldCrafterTransformer3DModel.from_pretrained(
paths["transformer"], torch_dtype=torch.bfloat16
)
patch_worldcrafter_transformer_ucpe(
transformer=transformer,
method="relray_absmap",
height=height,
width=width,
attn_compress=8,
adaptation_method="parallel",
)
camera_adapter = load_ucpe_camera_adapter_weights(transformer, paths["adapter"])
if (
camera_adapter["loaded_tensor_keys"]
!= camera_adapter["expected_tensor_keys"]
):
raise RuntimeError(f"camera adapter load is incomplete: {camera_adapter}")
adapter_dtypes = {
parameter.dtype
for block in transformer.blocks
for parameter in block.cam_self_attn.parameters()
}
if adapter_dtypes != {torch.float32}:
raise RuntimeError(f"camera adapter dtypes are invalid: {adapter_dtypes}")
if not enable_compile:
transformer = replace_rmsnorm_with_fp32(transformer)
transformer = replace_all_norms_with_flash_norms(transformer)
replace_rope_with_flash_rope()
resolved_backend = configure_attention(transformer, attention_backend)
pipeline = WorldCrafterPipeline(
tokenizer=AutoTokenizer.from_pretrained(paths["tokenizer"]),
text_encoder=UMT5EncoderModel.from_pretrained(
paths["text_encoder"], torch_dtype=torch.bfloat16
),
transformer=transformer,
vae=AutoencoderKLWan.from_pretrained(
paths["vae"], torch_dtype=torch.float32
),
scheduler=WorldCrafterScheduler.from_pretrained(paths["scheduler"]),
)
adapter_load = load_model_adapter(pipeline, paths["adapter"])
pipeline = pipeline.to(torch_device)
if enable_compile:
torch.backends.cudnn.benchmark = True
pipeline.text_encoder.compile(
mode="max-autotune-no-cudagraphs", dynamic=False
)
pipeline.vae.compile(mode="max-autotune-no-cudagraphs", dynamic=False)
pipeline.transformer.compile(
mode="max-autotune-no-cudagraphs", dynamic=False
)
return cls(
pipeline=pipeline,
memory_provider=memory_provider,
model_path=paths["root"],
device=torch_device,
attention_backend=resolved_backend,
adapter_load=adapter_load,
height=height,
width=width,
)
def generate(
self,
*,
mode: str,
camera_path: Path,
output_path: Path,
prompt: str,
negative_prompt: str,
image_path: Path | None = None,
num_chunks: int | None = None,
chunk_output_dir: Path | None = None,
state_output_dir: Path | None = None,
resume_from: Path | None = None,
on_chunk_saved: Callable[[int, Path], None] | None = None,
stop_after_chunk: int | None = None,
num_inference_steps: int | None = None,
guidance_scale: float | None = None,
seed: int = 42,
fps: int = 16,
image_noise_sigma_min: float = 0.111,
image_noise_sigma_max: float = 0.135,
camera_x_fov: float = 100.0,
camera_xi: float = 0.0,
local_camera_path: Path | None = None,
) -> InferenceResult:
if on_chunk_saved is not None and chunk_output_dir is None:
raise ValueError("on_chunk_saved requires chunk_output_dir")
is_fast = self.model_type == "fast"
if num_inference_steps is None:
num_inference_steps = 6 if is_fast else 50
if guidance_scale is None:
guidance_scale = 1.0 if is_fast else 5.0
if is_fast:
if guidance_scale != 1.0 or num_inference_steps != 6:
raise ValueError(
"Fast requires CFG=1 and six regular steps; the first T2V chunk uses twelve steps"
)
if resume_from is not None or state_output_dir is not None:
raise ValueError("Fast resume/state export is not yet validated")
self.pipeline.resident_branches.switch("equal")
self.pipeline.resident_branches.switches = 0
self.pipeline.stage_model_trace.clear()
self.pipeline.fast_inference_mode = mode
elif local_camera_path is not None:
raise ValueError("--local-camera-path is only used by fast inference")
if mode not in {"i2v", "t2v"}:
raise ValueError(f"unsupported mode: {mode}")
if not camera_path.is_file():
raise FileNotFoundError(camera_path)
if mode == "i2v":
if image_path is None or not image_path.is_file():
raise FileNotFoundError(image_path)
image = load_image(str(image_path)).resize((self.width, self.height))
else:
if image_path is not None:
raise ValueError(
"text-to-video inference does not accept an input image"
)
image = None
camera_c2w = load_camera(camera_path, num_chunks=num_chunks)
num_frames = int(camera_c2w.shape[1])
total_chunks = num_frames // CAMERA_CHUNK_FRAMES
camera = {
"c2w": camera_c2w,
"x_fov": torch.full(
(1,), camera_x_fov, device=self.device, dtype=torch.float32
),
"xi": torch.full((1,), camera_xi, device=self.device, dtype=torch.float32),
}
if is_fast:
if local_camera_path is not None:
local_c2w = load_camera(local_camera_path, num_chunks=total_chunks)
else:
from .ucpe.bridge import _relative_pose_chunk
local_c2w = torch.cat(
[
_relative_pose_chunk(
camera_c2w,
chunk_index=k,
window_num_frames=33,
device=self.device,
)
for k in range(total_chunks)
],
dim=1,
)
camera["pose"] = torch.as_tensor(
local_c2w, device=self.device, dtype=torch.float32
)
camera["c2w"] = torch.as_tensor(
camera_c2w, device=self.device, dtype=torch.float32
)
if chunk_output_dir is not None:
chunk_output_dir.mkdir(parents=True, exist_ok=True)
if resume_from is not None and chunk_output_dir is None:
raise ValueError(
"resuming requires --chunk-output-dir with completed prefix chunks"
)
if state_output_dir is not None:
state_output_dir.mkdir(parents=True, exist_ok=True)
final_chunk_index = total_chunks - 1
if stop_after_chunk is not None:
if stop_after_chunk < 0 or stop_after_chunk >= total_chunks:
raise ValueError("stop_after_chunk must identify a generated chunk")
final_chunk_index = int(stop_after_chunk)
run_contract = {
"mode": mode,
"camera_sha256": sha256(camera_path),
"image_sha256": sha256(image_path) if image_path is not None else None,
"prompt_sha256": hashlib.sha256(prompt.encode("utf-8")).hexdigest(),
"negative_prompt_sha256": hashlib.sha256(
negative_prompt.encode("utf-8")
).hexdigest(),
"num_chunks": total_chunks,
"num_inference_steps": num_inference_steps,
"guidance_scale": guidance_scale,
"seed": seed,
"image_noise_sigma_min": image_noise_sigma_min,
"image_noise_sigma_max": image_noise_sigma_max,
"camera_x_fov": camera_x_fov,
"camera_xi": camera_xi,
"repencoder_model_sha256": self.memory_provider.runtime.report[
"model_sha256"
],
}
resume_state: dict[str, object] | None = None
prior_history_selection: list[dict[str, object]] = []
if resume_from is not None:
if not resume_from.is_file():
raise FileNotFoundError(resume_from)
resume_state = torch.load(
resume_from, map_location="cpu", weights_only=False
)
if resume_state.get("run_contract") != run_contract:
raise ValueError(
"resume checkpoint does not match this inference run contract"
)
prior_history_selection = list(resume_state.get("history_selection", []))
if final_chunk_index < int(resume_state["next_chunk_index"]):
raise ValueError("stop_after_chunk precedes the resume point")
def save_chunk(chunk_index: int, current_video: torch.Tensor) -> None:
if chunk_output_dir is None:
return
frames = self.pipeline.video_processor.postprocess_video(
current_video, output_type="np"
)[0]
path = chunk_output_dir / f"chunk_{chunk_index:03d}_33f.mp4"
export_to_video(frames, str(path), fps=fps)
if on_chunk_saved is not None:
on_chunk_saved(chunk_index, path)
print(f"[worldcrafter] completed {path}", flush=True)
def save_state(chunk_index: int, state: dict[str, object]) -> None:
if state_output_dir is not None:
save_chunk_state(
chunk_index,
state,
state_output_dir=state_output_dir,
run_contract=run_contract,
history_selection=[
*prior_history_selection,
*[
record.to_jsonable()
for record in self.memory_provider.render_records
],
],
)
self.memory_provider.reset_sequence()
with torch.inference_mode():
frames = self.pipeline(
prompt=prompt,
negative_prompt=negative_prompt,
height=self.height,
width=self.width,
num_frames=num_frames,
num_inference_steps=num_inference_steps,
guidance_scale=guidance_scale,
generator=torch.Generator(device=self.device).manual_seed(seed),
memory_size=4,
history_sizes=[2, 1],
num_latent_frames_per_chunk=9,
keep_first_frame=True,
is_enable_stage2=is_fast,
pyramid_num_inference_steps_list=None if is_fast else [2, 2, 2],
is_skip_first_chunk=False,
is_amplify_first_chunk=False,
use_zero_init=False,
zero_steps=1,
image=image,
image_noise_sigma_min=image_noise_sigma_min,
image_noise_sigma_max=image_noise_sigma_max,
video=None,
video_noise_sigma_min=0.111,
video_noise_sigma_max=0.135,
camera_trajectory=camera,
memory_provider=self.memory_provider,
callback_on_chunk_end=(
save_chunk if chunk_output_dir is not None else None
),
callback_on_chunk_state=(
save_state if state_output_dir is not None else None
),
resume_state=resume_state,
stop_after_chunk=final_chunk_index,
).frames[0]
start_chunk = (
int(resume_state["next_chunk_index"]) if resume_state is not None else 0
)
expected_records = max(0, final_chunk_index - max(1, start_chunk) + 1)
if len(self.memory_provider.render_records) != expected_records:
raise RuntimeError(
f"expected {expected_records} RepEncoder calls, "
f"got {len(self.memory_provider.render_records)}"
)
output_path.parent.mkdir(parents=True, exist_ok=True)
if is_fast and len(frames) != (final_chunk_index + 1) * CAMERA_CHUNK_FRAMES:
raise RuntimeError("Fast output must preserve every decoded RGB frame")
if resume_state is None:
export_to_video(frames, str(output_path), fps=fps)
else:
assemble_resumed_video(output_path, chunk_output_dir, final_chunk_index)
history_selection = [
*prior_history_selection,
*[record.to_jsonable() for record in self.memory_provider.render_records],
]
summary = {
"format": "worldcrafter_inference_v2",
"mode": mode,
"model_path": str(self.model_path),
"image_path": str(image_path) if image_path is not None else None,
"image_sha256": sha256(image_path) if image_path is not None else None,
"camera_path": str(camera_path),
"camera_sha256": sha256(camera_path),
"camera_semantics": "global metric c2w; UCPE chunk-relative poses are derived internally",
"output_path": str(output_path),
"output_sha256": sha256(output_path),
"prompt": prompt,
"negative_prompt": negative_prompt,
"num_frames": num_frames,
"completed_through_chunk": final_chunk_index,
"fps": fps,
"seed": seed,
"num_inference_steps": num_inference_steps,
"guidance_scale": guidance_scale,
"attention_backend": self.attention_backend,
"adapter": self.adapter_load,
"repencoder_model_sha256": self.memory_provider.runtime.report[
"model_sha256"
],
"resumed_from": str(resume_from) if resume_from is not None else None,
"history_selection": history_selection,
}
if is_fast:
first_chunk_steps = 12 if mode == "t2v" else 6
expected_calls = first_chunk_steps + final_chunk_index * 6
if len(self.pipeline.stage_model_trace) != expected_calls:
raise RuntimeError(
"Fast forward count differs from the mode-specific DMD contract"
)
summary.update(
model_type="fast",
local_camera_path=str(local_camera_path),
local_camera_sha256=(
sha256(local_camera_path) if local_camera_path else None
),
first_chunk_inference_steps=first_chunk_steps,
first_chunk_routing="4+8" if mode == "t2v" else "5+1",
subsequent_chunk_routing="2+4" if mode == "t2v" else "5+1",
fast=self.fast_report,
stage_model_trace=self.pipeline.stage_model_trace,
resident_branch_switches=self.pipeline.resident_branches.switches,
)
summary_path = output_path.with_suffix(".json")
summary_path.write_text(json.dumps(summary, indent=2, sort_keys=True) + "\n")
print(f"[worldcrafter] saved {output_path}", flush=True)
return InferenceResult(output_path, summary_path, summary)
__all__ = ["InferenceResult", "WorldCrafter", "load_camera", "sha256"]