Download runtime/image3f_deploy.py from Ravenh97/roborender_image3f: direct link, hf CLI and curl.
- Browser
- Download file 26.3 kB
-
https://huggingface.co/Ravenh97/roborender_image3f/resolve/main/runtime/image3f_deploy.py
- Command line
-
hf download hf://Ravenh97/roborender_image3f/runtime/image3f_deploy.py
-
curl -L -o image3f_deploy.py https://huggingface.co/Ravenh97/roborender_image3f/resolve/main/runtime/image3f_deploy.py
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"]) | |
| 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()) | |