#!/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())