import dataclasses import datetime as dt import contextlib import json import logging import pathlib import re import time from typing import Any import numpy as np import torch import tyro from openpi.policies import libero_policy from openpi.policies import policy_config as _policy_config from openpi.policies import ur5_policy from openpi.shared import normalize as _normalize from serve_policy_from_checkpoint import load_train_config_from_checkpoint @dataclasses.dataclass class Args: checkpoint_dir: pathlib.Path config_name: str | None = None default_prompt: str | None = None pytorch_device: str | None = None num_denoise_steps: int = 2 lora_rank: int | None = None lora_alpha: float | None = None warmup_runs: int = 2 timed_runs: int = 10 profile_flops: bool = False example: str = "auto" train_log: pathlib.Path | None = None metrics_out_path: pathlib.Path | None = None table_out_path: pathlib.Path | None = None norm_stats_path: pathlib.Path | None = None model_label: str | None = None def main(args: Args) -> None: checkpoint_metadata = _load_checkpoint_metadata(args.checkpoint_dir) train_config = load_train_config_from_checkpoint( args.checkpoint_dir, config_name=args.config_name, lora_rank=args.lora_rank, lora_alpha=args.lora_alpha, ) norm_stats = _load_norm_stats_override(args.norm_stats_path) policy = _policy_config.create_trained_policy( train_config, args.checkpoint_dir, default_prompt=args.default_prompt, sample_kwargs={"num_steps": args.num_denoise_steps}, norm_stats=norm_stats, pytorch_device=args.pytorch_device, ) example = _make_example(args.example, train_config) last_result = None device = _policy_device(policy) if device is not None and device.type == "cuda": torch.cuda.reset_peak_memory_stats(device) with torch.inference_mode(): for _ in range(args.warmup_runs): last_result = policy.infer(example) _assert_finite_actions(last_result) if torch.cuda.is_available() and (args.pytorch_device is None or "cuda" in args.pytorch_device): torch.cuda.synchronize() model_infer_ms = [] wall_ms = [] timing_values = { "sample_total_ms": [], "vlm_prefix_ms": [], "denoising_ms": [], "denoising_ms_per_step": [], "denoising_steps": [], "other_ms": [], } with torch.inference_mode(): for _ in range(args.timed_runs): start_time = time.monotonic() result = policy.infer(example) if torch.cuda.is_available() and (args.pytorch_device is None or "cuda" in args.pytorch_device): torch.cuda.synchronize() wall_ms.append((time.monotonic() - start_time) * 1000) _assert_finite_actions(result) policy_timing = result.get("policy_timing", {}) model_infer_ms.append(float(policy_timing.get("infer_ms", float("nan")))) _append_timing_values(timing_values, policy_timing) last_result = result if last_result is None: raise RuntimeError("No inference runs were executed") actions = np.asarray(last_result["actions"]) action_horizon = actions.shape[0] model_infer_summary = _latency_summary(model_infer_ms) wall_summary = _latency_summary(wall_ms) model_metrics = _collect_model_metrics(policy, args.checkpoint_dir, checkpoint_metadata) training_metrics = _collect_training_metrics(checkpoint_metadata, args.train_log) approx_gflops = None flop_profile = None if args.profile_flops: flop_profile = _profile_policy_flops(policy, example, device) approx_gflops = flop_profile["total_gflops"] metrics = { "model_label": args.model_label or train_config.name, "config": train_config.name, "checkpoint": str(args.checkpoint_dir), "norm_stats_path": str(args.norm_stats_path.resolve()) if args.norm_stats_path else None, "global_step": checkpoint_metadata.get("global_step"), "warmup_runs": args.warmup_runs, "timed_runs": args.timed_runs, "num_denoise_steps": args.num_denoise_steps, "action_horizon": action_horizon, "actions_shape": list(actions.shape), "actions_mean": float(actions.mean()), "actions_std": float(actions.std()), "latency": { "model_infer_ms": model_infer_summary, "wall_ms": wall_summary, "model_infer_per_action_ms": _latency_summary([value / action_horizon for value in model_infer_ms]), "wall_per_action_ms": _latency_summary([value / action_horizon for value in wall_ms]), }, "latency_decomposition": _latency_decomposition_summary(timing_values), "frequency": _frequency_metrics(model_infer_summary, wall_summary, action_horizon), "model": model_metrics, "training": training_metrics, "approx_gflops": approx_gflops, "flop_profile": flop_profile, } metrics["table_row"] = _build_table_row(metrics) print(f"config={train_config.name}") print(f"checkpoint={args.checkpoint_dir}") print(f"warmup_runs={args.warmup_runs}") print(f"timed_runs={args.timed_runs}") print(f"num_denoise_steps={args.num_denoise_steps}") print(f"actions_shape={actions.shape}") print(f"actions_mean={actions.mean():.6f}") print(f"actions_std={actions.std():.6f}") _print_latency("model_infer_ms", model_infer_ms) _print_latency("wall_ms", wall_ms) _print_latency("model_infer_per_action_ms", [value / action_horizon for value in model_infer_ms]) _print_latency("wall_per_action_ms", [value / action_horizon for value in wall_ms]) decomposition = metrics["latency_decomposition"] if decomposition["available"]: if timing_values["vlm_prefix_ms"]: _print_latency("vlm_prefix_ms", timing_values["vlm_prefix_ms"]) if timing_values["denoising_ms"]: _print_latency("denoising_ms", timing_values["denoising_ms"]) if timing_values["denoising_ms_per_step"]: _print_latency("denoising_ms_per_step", timing_values["denoising_ms_per_step"]) print(f"model_query_hz={metrics['frequency']['model_query_hz']:.4f}") print(f"model_action_generation_hz={metrics['frequency']['model_action_generation_hz']:.4f}") print(f"wall_query_hz={metrics['frequency']['wall_query_hz']:.4f}") print(f"wall_action_generation_hz={metrics['frequency']['wall_action_generation_hz']:.4f}") print(f"model_total_params_b={model_metrics.get('model_total_params_b', float('nan')):.4f}") print(f"model_trainable_params_m={model_metrics.get('model_trainable_params_m', float('nan')):.2f}") print(f"checkpoint_size_gb={model_metrics.get('checkpoint_size_gb', float('nan')):.2f}") if training_metrics.get("training_time_hours") is not None: print(f"training_time_hours={training_metrics['training_time_hours']:.4f}") if approx_gflops is not None: print(f"approx_gflops={approx_gflops:.2f}") if flop_profile is not None: print(f"profiler_gflops={flop_profile['profiler_gflops']:.2f}") print(f"manual_conv2d_gflops={flop_profile['manual_conv2d_gflops']:.2f}") print(f"conv2d_modules_counted={flop_profile['manual_conv2d_modules_counted']}") print("table_row_markdown=" + _format_table_row(metrics["table_row"])) if args.metrics_out_path is not None: args.metrics_out_path.parent.mkdir(parents=True, exist_ok=True) args.metrics_out_path.write_text(json.dumps(metrics, indent=2) + "\n") print(f"metrics_out_path={args.metrics_out_path}") if args.table_out_path is not None: args.table_out_path.parent.mkdir(parents=True, exist_ok=True) args.table_out_path.write_text(_format_table(metrics["table_row"])) print(f"table_out_path={args.table_out_path}") def _assert_finite_actions(result: dict) -> None: actions = np.asarray(result["actions"]) if not np.isfinite(actions).all(): raise RuntimeError("Policy returned non-finite actions") def _make_example(example_name: str, train_config: Any) -> dict: example_name = example_name.lower() if example_name == "auto": config_text = f"{train_config.name} {type(train_config.data).__name__}".lower() example_name = "ur5" if "ur5" in config_text or "tube" in config_text or "pour" in config_text else "libero" if example_name == "ur5": return ur5_policy.make_ur5_example() if example_name == "libero": return libero_policy.make_libero_example() raise ValueError(f"Unknown example={example_name!r}. Expected auto, ur5, or libero.") def _print_latency(name: str, values: list[float]) -> None: summary = _latency_summary(values) print(f"{name}_mean={summary['mean']:.2f}") print(f"{name}_median={summary['median']:.2f}") print(f"{name}_p95={summary['p95']:.2f}") print(f"{name}_min={summary['min']:.2f}") print(f"{name}_max={summary['max']:.2f}") def _append_timing_values(timing_values: dict[str, list[float]], policy_timing: dict[str, Any]) -> None: for key in timing_values: value = _finite_float(policy_timing.get(key)) if value is not None: timing_values[key].append(value) def _latency_decomposition_summary(timing_values: dict[str, list[float]]) -> dict[str, Any]: sample_total = _latency_summary_or_none(timing_values["sample_total_ms"]) vlm_prefix = _latency_summary_or_none(timing_values["vlm_prefix_ms"]) denoising = _latency_summary_or_none(timing_values["denoising_ms"]) denoising_per_step = _latency_summary_or_none(timing_values["denoising_ms_per_step"]) denoising_steps = _latency_summary_or_none(timing_values["denoising_steps"]) other = _latency_summary_or_none(timing_values["other_ms"]) return { "available": bool(timing_values["sample_total_ms"]), "measurement": ( "Cached-observation policy benchmark. VLM prefix is image/language prefix embedding plus " "prefix transformer KV-cache construction. Denoising is the action denoising loop." ), "sample_total_ms": sample_total, "vlm_prefix_ms": vlm_prefix, "denoising_ms": denoising, "denoising_ms_per_step": denoising_per_step, "denoising_steps": denoising_steps, "other_ms": other, "fractions_of_mean_sample_total": { "vlm_prefix": _fraction(vlm_prefix.get("mean"), sample_total.get("mean")), "denoising": _fraction(denoising.get("mean"), sample_total.get("mean")), "other": _fraction(other.get("mean"), sample_total.get("mean")), }, } def _latency_summary_or_none(values: list[float]) -> dict[str, float | None]: if not values: return { "mean": None, "median": None, "p95": None, "min": None, "max": None, } return _latency_summary(values) def _finite_float(value: Any) -> float | None: try: number = float(value) except (TypeError, ValueError): return None return number if np.isfinite(number) else None def _fraction(numerator: Any, denominator: Any) -> float | None: num = _finite_float(numerator) den = _finite_float(denominator) if num is None or den is None or den <= 0: return None return float(num / den) def _latency_summary(values: list[float]) -> dict[str, float]: arr = np.asarray(values, dtype=np.float64) return { "mean": float(arr.mean()), "median": float(np.percentile(arr, 50)), "p95": float(np.percentile(arr, 95)), "min": float(arr.min()), "max": float(arr.max()), } def _frequency_metrics( model_infer_summary: dict[str, float], wall_summary: dict[str, float], action_horizon: int, ) -> dict[str, float]: model_query_hz = _hz_from_ms(model_infer_summary["median"]) wall_query_hz = _hz_from_ms(wall_summary["median"]) return { "model_query_hz": model_query_hz, "wall_query_hz": wall_query_hz, "model_action_generation_hz": model_query_hz * action_horizon, "wall_action_generation_hz": wall_query_hz * action_horizon, } def _hz_from_ms(milliseconds: float) -> float: return 1000.0 / milliseconds if milliseconds > 0 else float("nan") def _load_norm_stats_override(path: pathlib.Path | None): if path is None: return None path = pathlib.Path(path).expanduser().resolve() if path.is_file(): logging.info("Loaded explicit norm stats from %s", path) return _normalize.deserialize_json(path.read_text(encoding="utf-8")) logging.info("Loaded explicit norm stats from %s", path / "norm_stats.json") return _normalize.load(path) def _load_checkpoint_metadata(checkpoint_dir: pathlib.Path) -> dict[str, Any]: metadata_path = checkpoint_dir / "metadata.pt" if not metadata_path.exists(): return {} try: return torch.load(metadata_path, map_location="cpu", weights_only=False) except TypeError: return torch.load(metadata_path, map_location="cpu") def _policy_device(policy) -> torch.device | None: device_name = getattr(policy, "_pytorch_device", None) if device_name is not None: return torch.device(device_name) model = getattr(policy, "_model", None) if model is None: return None try: return next(model.parameters()).device except StopIteration: return None def _collect_model_metrics(policy, checkpoint_dir: pathlib.Path, metadata: dict[str, Any]) -> dict[str, float | int | None]: model = getattr(policy, "_model", None) metrics = dict(metadata.get("efficiency_metrics") or {}) if model is not None: if hasattr(model, "configure_trainable_parameters"): model.configure_trainable_parameters() params = list(model.parameters()) total_params = sum(param.numel() for param in params) trainable_params = sum(param.numel() for param in params if param.requires_grad) metrics.update( { "model_total_params": total_params, "model_trainable_params": trainable_params, "model_trainable_pct": trainable_params / total_params * 100 if total_params else 0.0, "model_total_params_b": total_params / 1e9, "model_trainable_params_m": trainable_params / 1e6, } ) weight_path = checkpoint_dir / "model.safetensors" if weight_path.exists(): checkpoint_size_bytes = weight_path.stat().st_size metrics["checkpoint_size_bytes"] = checkpoint_size_bytes metrics["checkpoint_size_gb"] = checkpoint_size_bytes / 1024**3 device = _policy_device(policy) if device is not None and device.type == "cuda" and torch.cuda.is_available(): metrics["cuda_peak_allocated_gb_inference"] = torch.cuda.max_memory_allocated(device) / 1024**3 metrics["cuda_peak_reserved_gb_inference"] = torch.cuda.max_memory_reserved(device) / 1024**3 return metrics def _collect_training_metrics(metadata: dict[str, Any], train_log: pathlib.Path | None) -> dict[str, Any]: efficiency_metrics = metadata.get("efficiency_metrics") or {} training_time_hours = ( efficiency_metrics.get("elapsed_training_hours") or efficiency_metrics.get("elapsed_after_first_step_hours") or efficiency_metrics.get("elapsed_run_hours") ) metrics: dict[str, Any] = { "training_time_hours": training_time_hours, "source": "checkpoint_metadata" if training_time_hours is not None else None, } if train_log is not None: parsed = _parse_training_time_from_log(train_log, metadata.get("global_step")) metrics.update(parsed) if parsed.get("training_time_hours") is not None: metrics["training_time_hours"] = parsed["training_time_hours"] metrics["source"] = "train_log" return metrics def _parse_training_time_from_log(train_log: pathlib.Path, global_step: int | None) -> dict[str, Any]: if not train_log.exists(): raise FileNotFoundError(f"Training log does not exist: {train_log}") created = None train_config = None first_step = None final_step = None final_save = None step_pattern = re.compile(rf"step={global_step}\s") if global_step is not None else None for raw_line in train_log.read_text(errors="replace").replace("\r", "\n").splitlines(): timestamp = _parse_log_time(raw_line) if timestamp is None: continue if "Created experiment checkpoint directory" in raw_line: created = timestamp if "Training config:" in raw_line: train_config = timestamp if re.search(r"step=1\s", raw_line): first_step = timestamp if step_pattern is not None and step_pattern.search(raw_line): final_step = timestamp if global_step is not None and f"Saved checkpoint at step {global_step}" in raw_line: final_save = timestamp end_time = final_save or final_step start_time = train_config or first_step or created training_time_hours = None if start_time is None or end_time is None else _elapsed_hours(start_time, end_time) return { "training_time_hours": training_time_hours, "created_to_final_save_hours": None if created is None or final_save is None else _elapsed_hours(created, final_save), "train_config_to_final_save_hours": ( None if train_config is None or final_save is None else _elapsed_hours(train_config, final_save) ), "first_step_to_final_save_hours": None if first_step is None or final_save is None else _elapsed_hours(first_step, final_save), "first_step_to_final_step_hours": None if first_step is None or final_step is None else _elapsed_hours(first_step, final_step), "log_path": str(train_log), } def _parse_log_time(line: str) -> dt.datetime | None: match = re.search(r"(\d{2}:\d{2}:\d{2}\.\d{3})", line) if match is None: return None return dt.datetime.strptime(match.group(1), "%H:%M:%S.%f") def _elapsed_hours(start: dt.datetime, end: dt.datetime) -> float: while end < start: end += dt.timedelta(days=1) return (end - start).total_seconds() / 3600 def _profile_policy_flops(policy, example: dict, device: torch.device | None) -> dict[str, Any]: activities = [torch.profiler.ProfilerActivity.CPU] if device is not None and device.type == "cuda" and torch.cuda.is_available(): activities.append(torch.profiler.ProfilerActivity.CUDA) torch.cuda.synchronize(device) conv2d_counter = _Conv2dFlopCounter(policy) with torch.inference_mode(): with conv2d_counter: with torch.profiler.profile(activities=activities, with_flops=True) as profiler: result = policy.infer(example) _assert_finite_actions(result) if device is not None and device.type == "cuda" and torch.cuda.is_available(): torch.cuda.synchronize(device) profiler_flops = sum((getattr(event, "flops", 0) or 0) for event in profiler.key_averages()) profiler_conv2d_flops = sum( (getattr(event, "flops", 0) or 0) for event in profiler.key_averages() if "conv2d" in getattr(event, "key", "") ) manual_conv2d_flops = conv2d_counter.total_flops # PyTorch currently warns that some aten::conv2d FLOPs cannot be computed, # and reports those as zero. Add our module-hook Conv2d count only when the # profiler did not already account for Conv2d FLOPs to avoid double counting # on future PyTorch versions. added_manual_conv2d_flops = manual_conv2d_flops if profiler_conv2d_flops == 0 else 0 total_flops = profiler_flops + added_manual_conv2d_flops return { "total_flops": float(total_flops), "total_gflops": float(total_flops) / 1e9, "profiler_flops": float(profiler_flops), "profiler_gflops": float(profiler_flops) / 1e9, "profiler_conv2d_flops": float(profiler_conv2d_flops), "profiler_conv2d_gflops": float(profiler_conv2d_flops) / 1e9, "manual_conv2d_flops": float(manual_conv2d_flops), "manual_conv2d_gflops": float(manual_conv2d_flops) / 1e9, "manual_conv2d_flops_added": float(added_manual_conv2d_flops), "manual_conv2d_gflops_added": float(added_manual_conv2d_flops) / 1e9, "manual_conv2d_modules_counted": conv2d_counter.modules_counted, "flop_convention": "2 FLOPs per multiply-add; bias additions are not counted.", "note": ( "Total combines PyTorch profiler FLOPs with manual Conv2d FLOPs only " "when profiler Conv2d FLOPs are zero." ), } class _Conv2dFlopCounter(contextlib.AbstractContextManager): def __init__(self, policy) -> None: self.model = getattr(policy, "_model", None) self.handles: list[Any] = [] self.total_flops = 0.0 self.modules_counted = 0 def __enter__(self): if self.model is None: return self for module in self.model.modules(): if isinstance(module, torch.nn.Conv2d): self.handles.append(module.register_forward_hook(self._hook)) return self def __exit__(self, exc_type, exc_value, traceback) -> bool: for handle in self.handles: handle.remove() self.handles.clear() return False def _hook(self, module: torch.nn.Conv2d, inputs: tuple[Any, ...], output: Any) -> None: if not inputs: return x = inputs[0] if not torch.is_tensor(x) or not torch.is_tensor(output) or output.ndim < 4: return batch = int(output.shape[0]) out_channels = int(output.shape[1]) out_h = int(output.shape[2]) out_w = int(output.shape[3]) kernel_h, kernel_w = module.kernel_size in_channels = int(module.in_channels) groups = int(module.groups) macs_per_output = (in_channels // groups) * int(kernel_h) * int(kernel_w) self.total_flops += float(batch * out_channels * out_h * out_w * macs_per_output * 2) self.modules_counted += 1 def _build_table_row(metrics: dict[str, Any]) -> dict[str, Any]: model_metrics = metrics["model"] training_metrics = metrics["training"] return { "Model": metrics["model_label"], "Model Size": f"{model_metrics.get('model_total_params_b', float('nan')):.2f}B", "Trainable Params": f"{model_metrics.get('model_trainable_params_m', float('nan')):.2f}M", "Training Time (hours)": _format_optional(training_metrics.get("training_time_hours"), precision=3), "GFLOPs": _format_optional(metrics.get("approx_gflops"), precision=2), "Inference Speed (ms)": f"{metrics['latency']['model_infer_ms']['median']:.2f}", "Inference Speed (Hz)": f"{metrics['frequency']['model_query_hz']:.3f}", "Action Gen. (Hz)": f"{metrics['frequency']['model_action_generation_hz']:.3f}", "Checkpoint Size": f"{model_metrics.get('checkpoint_size_gb', float('nan')):.2f}GB", } def _format_optional(value: Any, *, precision: int) -> str: if value is None: return "n/a" return f"{float(value):.{precision}f}" def _format_table_row(row: dict[str, Any]) -> str: return "| " + " | ".join(str(value) for value in row.values()) + " |" def _format_table(row: dict[str, Any]) -> str: header = "| " + " | ".join(row.keys()) + " |\n" separator = "| " + " | ".join("---" for _ in row) + " |\n" return header + separator + _format_table_row(row) + "\n" if __name__ == "__main__": logging.basicConfig(level=logging.INFO, force=True) main(tyro.cli(Args))