roborender_image3f / runtime /image3f_deploy.py
Ravenh97's picture
Image3F rgb030-depth030 step-5000: adapter + deployment source
dd0a8f2 verified
Raw History Blame Contribute Delete
26.3 kB
#!/usr/bin/env python3
"""Persistent fixed-shape deployment runtime for RoboRender Image3F.
The TensorRT plan is deliberately not portable. ``bin/build_engine.sh``
creates a plan for the GPU on which this runtime will execute and records the
compatibility tuple under ``assets/generated/engines/current``.
"""
from __future__ import annotations
import argparse
import hashlib
import json
import os
import statistics
import subprocess
import sys
import threading
import time
from dataclasses import dataclass
from pathlib import Path
from typing import Any, Iterable
from PIL import Image
import torch
ROOT = Path(__file__).resolve().parents[1]
if str(ROOT) not in sys.path:
sys.path.insert(0, str(ROOT))
from runtime.release_contract import ( # noqa: E402
adapter_path,
current_engine_dir,
export_lineage,
foundation_model_dir,
generated_onnx_dir,
load_config,
resolve_model_root,
sha256,
tokenizer_path,
)
CONFIG = load_config()
ASSETS = ROOT / "assets"
VENDOR = ROOT / "vendor" / "raven-image3f-joint"
CONTROL = VENDOR / "control"
for candidate in (str(CONTROL), str(VENDOR)):
if candidate not in sys.path:
sys.path.insert(0, candidate)
os.environ.setdefault("DIFFSYNTH_SKIP_DOWNLOAD", "true")
from benchmark_image3f_latency import setup_multiview # noqa: E402
from diffsynth.models.wan_video_dit import ( # noqa: E402
set_attention_implementation,
set_rope_implementation,
set_sdpa_backend,
)
from diffsynth.pipelines.wan_video import ModelConfig, WanVideoPipeline # noqa: E402
from image3f_trt_runtime import install_trt_block_stack # noqa: E402
VIEW_NAMES = ("ext1", "ext2", "wrist")
WIDTH = 416
PER_VIEW_HEIGHT = 240
LATENT_SHAPE = (1, 16, 1, 90, 52)
TIMESTEP_SCHEDULE = tuple(CONFIG["generation_contract"]["timestep_schedule"])
@dataclass(frozen=True)
class Image3FRequest:
prompt: str
depth: tuple[Image.Image, Image.Image, Image.Image]
mask: tuple[Image.Image, Image.Image, Image.Image]
previous_rgb: tuple[Image.Image, Image.Image, Image.Image]
seed: int = 42
def _read_rgb(path: str | Path) -> Image.Image:
image = Image.open(path).convert("RGB")
if image.size != (WIDTH, PER_VIEW_HEIGHT):
raise ValueError(
f"{path}: expected {WIDTH}x{PER_VIEW_HEIGHT}, found "
f"{image.width}x{image.height}; resize explicitly upstream"
)
return image
def request_from_mapping(value: dict[str, Any]) -> Image3FRequest:
def three_paths(key: str) -> tuple[Image.Image, Image.Image, Image.Image]:
paths = value.get(key)
if not isinstance(paths, list) or len(paths) != 3:
raise ValueError(f"{key} must be a three-element path list in ext1/ext2/wrist order")
return tuple(_read_rgb(path) for path in paths) # type: ignore[return-value]
prompt = value.get("prompt")
if not isinstance(prompt, str) or not prompt.strip():
raise ValueError("prompt must be a non-empty string")
return Image3FRequest(
prompt=prompt,
depth=three_paths("depth"),
mask=three_paths("mask"),
previous_rgb=three_paths("previous_rgb"),
seed=int(value.get("seed", 42)),
)
def load_request_json(path: str | Path) -> Image3FRequest:
source = Path(path).resolve()
value = json.loads(source.read_text(encoding="utf-8"))
for key in ("depth", "mask", "previous_rgb"):
if isinstance(value.get(key), list):
value[key] = [
str((source.parent / item).resolve()) if not Path(item).is_absolute() else item
for item in value[key]
]
return request_from_mapping(value)
def _validate_request(request: Image3FRequest) -> None:
if not isinstance(request.prompt, str) or not request.prompt.strip():
raise ValueError("prompt must be a non-empty string")
for field_name in ("depth", "mask", "previous_rgb"):
images = getattr(request, field_name)
if len(images) != 3:
raise ValueError(
f"{field_name} must contain ext1/ext2/wrist in that order"
)
for view_name, image in zip(VIEW_NAMES, images):
if not isinstance(image, Image.Image):
raise TypeError(f"{field_name}.{view_name} must be a PIL image")
if image.mode != "RGB" or image.size != (WIDTH, PER_VIEW_HEIGHT):
raise ValueError(
f"{field_name}.{view_name} must be RGB {WIDTH}x{PER_VIEW_HEIGHT}; "
f"found {image.mode} {image.width}x{image.height}"
)
def _offload_pytorch_block_fallback(pipe: WanVideoPipeline) -> int:
"""Move the duplicate block weights to CPU before loading the TRT plan."""
blocks = getattr(pipe.dit, "blocks", None)
if blocks is None:
raise RuntimeError("pipe.dit.blocks is unavailable")
parameter_bytes = sum(p.numel() * p.element_size() for p in blocks.parameters())
blocks.to(device="cpu")
torch.cuda.empty_cache()
return parameter_bytes
def _resolve_engine(
explicit: str | None,
*,
config: dict[str, Any],
model_root: Path,
) -> tuple[Path, dict[str, Any]]:
path = (
Path(explicit).resolve()
if explicit
else current_engine_dir(config)
/ "image3f_dit_block_stack_reference_bf16.engine"
)
if not path.is_file():
raise FileNotFoundError(
f"No target-compatible TensorRT engine at {path}. Run bin/build_engine.sh on this GPU."
)
manifest_path = path.parent / "engine_build_manifest.json"
if not manifest_path.is_file():
raise FileNotFoundError(f"Missing engine compatibility manifest: {manifest_path}")
manifest = json.loads(manifest_path.read_text(encoding="utf-8"))
engine_row = manifest.get("engine", {})
if engine_row.get("file") != path.name:
raise RuntimeError(
f"Engine manifest names {engine_row.get('file')!r}, but selected file is {path.name!r}"
)
expected_sha = engine_row.get("sha256")
if not expected_sha:
raise RuntimeError("Engine manifest does not contain a SHA-256")
digest = hashlib.sha256()
with path.open("rb") as handle:
for block in iter(lambda: handle.read(16 * 1024 * 1024), b""):
digest.update(block)
if digest.hexdigest() != expected_sha:
raise RuntimeError(f"TensorRT engine SHA-256 mismatch: {path}")
import tensorrt as trt
capability = ".".join(map(str, torch.cuda.get_device_capability(0)))
properties = torch.cuda.get_device_properties(0)
driver = subprocess.check_output(
["nvidia-smi", "--query-gpu=driver_version", "--format=csv,noheader"],
text=True,
).splitlines()[0].strip()
errors = []
if manifest.get("compute_capability") != capability:
errors.append(
f"compute capability {manifest.get('compute_capability')} != current {capability}"
)
if manifest.get("gpu") != properties.name:
errors.append(f"GPU {manifest.get('gpu')!r} != current {properties.name!r}")
if manifest.get("tensorrt") != trt.__version__:
errors.append(f"TensorRT {manifest.get('tensorrt')} != current {trt.__version__}")
if manifest.get("torch") != torch.__version__:
errors.append(f"PyTorch {manifest.get('torch')} != current {torch.__version__}")
if manifest.get("cuda") != torch.version.cuda:
errors.append(f"CUDA {manifest.get('cuda')} != current {torch.version.cuda}")
if manifest.get("driver") != driver:
errors.append(f"driver {manifest.get('driver')} != current {driver}")
lineage = export_lineage(model_root, config=config)
export_manifest_path = (
generated_onnx_dir(model_root, config) / "export_manifest.json"
)
if not export_manifest_path.is_file():
errors.append(f"missing ONNX export manifest {export_manifest_path}")
elif manifest.get("export_manifest_sha256") != sha256(export_manifest_path):
errors.append("engine was built from a different ONNX export manifest")
if manifest.get("adapter_sha256") != config["adapter"]["sha256"]:
errors.append("engine adapter SHA does not match this release")
if manifest.get("export_key") != lineage["export_key"]:
errors.append("engine ONNX lineage does not match this release")
if errors:
raise RuntimeError(
"TensorRT engine is not compatible with this runtime ("
+ "; ".join(errors)
+ "). Rebuild it locally with bin/build_engine.sh."
)
return path, manifest
class Image3FSession:
"""Reusable, serialized inference session for a robot control process.
Keep one instance alive. Model load, TensorRT deserialization, and the exact
prompt embedding cache are intentionally amortized across requests.
"""
def __init__(
self,
*,
model_root: str | Path | None = None,
engine_path: str | None = None,
memory_mode: str = "auto",
tea_cache_threshold: float | None = None,
cfg_scale: float | None = None,
vae_mode: str = "untiled",
):
if memory_mode not in {"auto", "resident", "24gb"}:
raise ValueError("memory_mode must be auto, resident, or 24gb")
if vae_mode not in {"untiled", "tiled"}:
raise ValueError("vae_mode must be untiled or tiled")
if not torch.cuda.is_available():
raise RuntimeError("CUDA is required")
if not torch.cuda.is_bf16_supported():
raise RuntimeError("This BF16 release requires a BF16-capable NVIDIA GPU")
self.config = load_config()
self.model_base = resolve_model_root(model_root, self.config)
self.memory_mode = memory_mode
self.tea_cache_threshold = float(
self.config["acceleration"]["tea_cache"]["threshold"]
if tea_cache_threshold is None
else tea_cache_threshold
)
self.cfg_scale = float(
self.config["generation_contract"]["cfg_scale"]
if cfg_scale is None
else cfg_scale
)
self.vae_mode = vae_mode
self._lock = threading.Lock()
self.request_count = 0
self.load_started_s = time.perf_counter()
total_vram_gib = torch.cuda.get_device_properties(0).total_memory / 1024**3
self.cpu_offload_enabled = memory_mode == "24gb" or (
memory_mode == "auto" and total_vram_gib <= 32.0
)
self.engine_path, self.engine_manifest = _resolve_engine(
engine_path,
config=self.config,
model_root=self.model_base,
)
set_rope_implementation("real")
set_attention_implementation("sdpa")
set_sdpa_backend("flash")
# The 10.6 GiB T5 encoder is only needed on a prompt-cache miss. On a
# 24 GiB card, manage that model from CPU while leaving the VAE and
# CLIP image encoder resident; moving those per frame would add PCIe
# latency to the real-time path.
managed_text_encoder: dict[str, Any] = {}
if self.cpu_offload_enabled:
managed_text_encoder = {
"offload_device": "cpu",
"offload_dtype": torch.bfloat16,
"onload_device": "cuda",
"onload_dtype": torch.bfloat16,
"preparing_device": "cuda",
"preparing_dtype": torch.bfloat16,
"computation_device": "cuda",
"computation_dtype": torch.bfloat16,
}
model_root = foundation_model_dir(self.model_base, self.config)
tokenizer = tokenizer_path(self.config)
checkpoint = adapter_path(self.config)
for required in (model_root, tokenizer, checkpoint):
if not required.exists():
raise FileNotFoundError(required)
if sha256(checkpoint) != self.config["adapter"]["sha256"]:
raise RuntimeError("adapter SHA-256 does not match config/deployment.json")
os.environ["DIFFSYNTH_MODEL_BASE_PATH"] = str(self.model_base)
self.pipe = WanVideoPipeline.from_pretrained(
torch_dtype=torch.bfloat16,
device="cuda",
model_configs=[
ModelConfig(
path=str(model_root / "diffusion_pytorch_model.safetensors")
),
ModelConfig(
path=str(model_root / "models_t5_umt5-xxl-enc-bf16.pth"),
**managed_text_encoder,
),
ModelConfig(path=str(model_root / "Wan2.1_VAE.pth")),
ModelConfig(
path=str(model_root / "models_clip_open-clip-xlm-roberta-large-vit-huge-14.pth"),
),
],
tokenizer_config=ModelConfig(path=str(tokenizer)),
redirect_common_files=False,
)
self.pipe.load_lora(
self.pipe.dit,
str(checkpoint),
alpha=float(self.config["adapter"]["inference_alpha"]),
)
setup_multiview(self.pipe)
self.offloaded_block_bytes = _offload_pytorch_block_fallback(self.pipe)
self.dispatcher = install_trt_block_stack(
self.pipe, [str(self.engine_path)], use_cuda_graph=False
)
self.load_seconds = time.perf_counter() - self.load_started_s
def describe(self) -> dict[str, Any]:
props = torch.cuda.get_device_properties(0)
prompt_cache_stats = dict(
getattr(self.pipe, "_image3f_prompt_cache_stats", {})
)
return {
"gpu": props.name,
"compute_capability": ".".join(map(str, torch.cuda.get_device_capability(0))),
"total_vram_gib": props.total_memory / 1024**3,
"memory_mode": self.memory_mode,
"cpu_offload_enabled": self.cpu_offload_enabled,
"offloaded_pytorch_block_gib": self.offloaded_block_bytes / 1024**3,
"engine_path": str(self.engine_path),
"load_seconds": self.load_seconds,
"steps": 5,
"cfg_scale": self.cfg_scale,
"tea_cache_threshold": self.tea_cache_threshold,
"vae_mode": self.vae_mode,
"cuda_graph": False,
"prompt_cache": "persistent exact-prompt",
"prompt_cache_stats": prompt_cache_stats,
}
def generate(self, request: Image3FRequest, *, output_type: str = "latent") -> tuple[Any, dict[str, Any]]:
if output_type not in {"latent", "rgb"}:
raise ValueError("output_type must be latent or rgb")
_validate_request(request)
trace: dict[str, Any] = {}
with self._lock:
trt_before = sum(engine.invocations for engine in self.dispatcher.engines)
torch.cuda.reset_peak_memory_stats()
torch.cuda.synchronize()
started = time.perf_counter()
output = self.pipe(
prompt=request.prompt,
negative_prompt="",
control_video_views=[[image] for image in request.depth],
control_video_2_views=[[image] for image in request.mask],
reference_image_views=list(request.previous_rgb),
height=PER_VIEW_HEIGHT * 3,
width=WIDTH,
num_frames=1,
num_views=3,
seed=request.seed,
cfg_scale=self.cfg_scale,
num_inference_steps=5,
timestep_schedule=list(TIMESTEP_SCHEDULE),
tea_cache_l1_thresh=self.tea_cache_threshold,
tea_cache_model_id="Image3FJoint-5step",
enable_static_cache=True,
enable_prompt_cache=True,
benchmark_trace=trace,
return_latents=output_type == "latent",
tiled=self.vae_mode == "tiled",
progress_bar_cmd=lambda values: values,
)
torch.cuda.synchronize()
elapsed_ms = (time.perf_counter() - started) * 1000.0
peak_gib = torch.cuda.max_memory_allocated() / 1024**3
current_gib = torch.cuda.memory_allocated() / 1024**3
trt_after = sum(engine.invocations for engine in self.dispatcher.engines)
trt_invocations = trt_after - trt_before
self.request_count += 1
process_resident_gib = _current_process_gpu_memory_gib()
estimated_process_peak_gib = (
None
if process_resident_gib is None
else process_resident_gib + max(0.0, peak_gib - current_gib)
)
traced_block_calls = trace.get("tea_cache_block_stack_calls")
if traced_block_calls is None or int(traced_block_calls) != trt_invocations:
raise RuntimeError(
"TensorRT/TeaCache execution contract mismatch: "
f"trace={traced_block_calls}, TensorRT={trt_invocations}"
)
if output_type == "latent":
if not torch.is_tensor(output) or tuple(output.shape) != LATENT_SHAPE:
raise RuntimeError(f"Unexpected latent output: {type(output)!r}, {getattr(output, 'shape', None)}")
finite = bool(torch.isfinite(output).all().item())
else:
if not isinstance(output, list) or len(output) != 1 or output[0].size != (WIDTH, PER_VIEW_HEIGHT * 3):
raise RuntimeError(f"Unexpected RGB output contract: {type(output)!r}")
finite = True
timing = {
"request_index": self.request_count - 1,
"output_type": output_type,
"pipeline_ms": elapsed_ms,
"unit_preprocessing_ms": _seconds_to_ms(trace.get("unit_preprocessing_s")),
"denoising_loop_ms": _seconds_to_ms(trace.get("denoising_loop_s")),
"post_units_ms": _seconds_to_ms(trace.get("post_units_s")),
"vae_decode_ms": _seconds_to_ms(trace.get("vae_decode_s")),
"tea_cache_block_stack_calls": trace.get("tea_cache_block_stack_calls"),
"tea_cache_skips": trace.get("tea_cache_skips"),
"trt_invocations": trt_invocations,
"peak_allocated_gib": peak_gib,
"current_allocated_gib": current_gib,
"process_gpu_resident_gib": process_resident_gib,
"estimated_process_gpu_peak_gib": estimated_process_peak_gib,
"finite": finite,
}
return output, timing
def _seconds_to_ms(value: Any) -> float | None:
return None if value is None else float(value) * 1000.0
def _current_process_gpu_memory_gib() -> float | None:
"""Return this process's post-request GPU residency for the active device.
PyTorch allocator statistics do not include TensorRT or CUDA-driver
allocations. ``nvidia-smi`` supplies that missing process-level view. A
failure is non-fatal because some container profiles hide accounting data.
"""
try:
active_uuid = (
str(torch.cuda.get_device_properties(0).uuid)
.lower()
.removeprefix("gpu-")
)
output = subprocess.check_output(
[
"nvidia-smi",
"--query-compute-apps=gpu_uuid,pid,used_gpu_memory",
"--format=csv,noheader,nounits",
],
text=True,
stderr=subprocess.DEVNULL,
)
used_mib = 0.0
matched = False
for line in output.splitlines():
fields = [field.strip() for field in line.split(",")]
if len(fields) != 3:
continue
gpu_uuid = fields[0].lower().removeprefix("gpu-")
if gpu_uuid == active_uuid and int(fields[1]) == os.getpid():
matched = True
used_mib += float(fields[2])
return used_mib / 1024.0 if matched else None
except (FileNotFoundError, OSError, ValueError, subprocess.SubprocessError):
return None
def _horizontal_views(stacked: Image.Image) -> Image.Image:
result = Image.new("RGB", (WIDTH * 3, PER_VIEW_HEIGHT))
for index in range(3):
crop = stacked.crop((0, index * PER_VIEW_HEIGHT, WIDTH, (index + 1) * PER_VIEW_HEIGHT))
result.paste(crop, (index * WIDTH, 0))
return result
def save_output(output: Any, output_type: str, output_dir: str | Path, stem: str) -> list[str]:
destination = Path(output_dir)
destination.mkdir(parents=True, exist_ok=True)
written: list[str] = []
if output_type == "latent":
from safetensors.torch import save_file
path = destination / f"{stem}.safetensors"
save_file(
{"latents": output.detach().cpu().contiguous()},
str(path),
metadata={"axis_order": "N,C,T,H,W", "vae_decoded": "false"},
)
written.append(str(path))
else:
stacked = output[0]
stacked_path = destination / f"{stem}_stacked.png"
horizontal_path = destination / f"{stem}_horizontal.png"
stacked.save(stacked_path)
_horizontal_views(stacked).save(horizontal_path)
written.extend((str(stacked_path), str(horizontal_path)))
return written
def _summary(records: Iterable[dict[str, Any]]) -> dict[str, float]:
values = [float(row["pipeline_ms"]) for row in records]
ordered = sorted(values)
return {
"count": len(values),
"mean_ms": statistics.mean(values),
"median_ms": statistics.median(values),
"p90_ms": ordered[max(0, int(0.9 * len(ordered) + 0.999999) - 1)],
"min_ms": min(values),
"max_ms": max(values),
}
def _serve_stdio(
session: Image3FSession,
default_output_type: str,
default_output_dir: Path,
warmup_request: Image3FRequest | None = None,
) -> int:
warmup_timing = None
if warmup_request is not None:
_, warmup_timing = session.generate(
warmup_request, output_type=default_output_type
)
print(
json.dumps(
{
"event": "ready",
"session": session.describe(),
"warmup_timing": warmup_timing,
}
),
flush=True,
)
for line in sys.stdin:
try:
spec = json.loads(line)
if spec.get("command") == "quit":
print(json.dumps({"event": "bye"}), flush=True)
return 0
request = request_from_mapping(spec)
output_type = spec.get("output_type", default_output_type)
output, timing = session.generate(request, output_type=output_type)
paths = save_output(
output,
output_type,
spec.get("output_dir", str(default_output_dir)),
spec.get("stem", f"request_{session.request_count - 1:06d}"),
)
print(json.dumps({"event": "result", "timing": timing, "outputs": paths}), flush=True)
except Exception as exc:
print(json.dumps({"event": "error", "error": f"{type(exc).__name__}: {exc}"}), flush=True)
return 0
def main() -> int:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument(
"--model-root",
help="External foundation-model root; defaults to ROBORENDER_MODEL_BASE",
)
parser.add_argument("--engine")
parser.add_argument("--memory-mode", choices=("auto", "resident", "24gb"), default="auto")
parser.add_argument("--tea-cache-threshold", type=float)
parser.add_argument("--cfg-scale", type=float)
parser.add_argument(
"--vae-mode",
choices=("untiled", "tiled"),
default="untiled",
help="Use the faster full-frame VAE or the lower-activation tiled fallback",
)
subparsers = parser.add_subparsers(dest="command", required=True)
run = subparsers.add_parser("run", help="Generate from one JSON request")
run.add_argument("--request", required=True)
run.add_argument("--output-type", choices=("latent", "rgb"), default="latent")
run.add_argument("--output-dir", default=str(ROOT / "outputs"))
run.add_argument("--warmup", type=int, default=1)
run.add_argument("--iterations", type=int, default=1)
serve = subparsers.add_parser("serve-stdio", help="Persistent JSON-lines inference service")
serve.add_argument("--output-type", choices=("latent", "rgb"), default="latent")
serve.add_argument("--output-dir", default=str(ROOT / "outputs"))
serve.add_argument(
"--warmup-request",
help="Prime prompt/static/VAE caches before emitting the ready event",
)
args = parser.parse_args()
session = Image3FSession(
model_root=args.model_root,
engine_path=args.engine,
memory_mode=args.memory_mode,
tea_cache_threshold=args.tea_cache_threshold,
cfg_scale=args.cfg_scale,
vae_mode=args.vae_mode,
)
if args.command == "serve-stdio":
warmup_request = (
load_request_json(args.warmup_request) if args.warmup_request else None
)
return _serve_stdio(
session,
args.output_type,
Path(args.output_dir),
warmup_request,
)
request = load_request_json(args.request)
if args.warmup < 0 or args.iterations < 1:
parser.error("--warmup must be non-negative and --iterations must be positive")
for _ in range(args.warmup):
session.generate(request, output_type=args.output_type)
records = []
last_output = None
for _ in range(args.iterations):
last_output, timing = session.generate(request, output_type=args.output_type)
records.append(timing)
print(json.dumps({"event": "timing", **timing}), flush=True)
outputs = save_output(last_output, args.output_type, args.output_dir, "image3f")
report = {
"session": session.describe(),
"summary": _summary(records),
"records": records,
"outputs": outputs,
}
report_path = Path(args.output_dir) / "run_manifest.json"
report_path.write_text(json.dumps(report, indent=2) + "\n", encoding="utf-8")
print(json.dumps({"event": "complete", **report}, indent=2), flush=True)
return 0
if __name__ == "__main__":
raise SystemExit(main())