File size: 15,345 Bytes
e479c46 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 | # SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Single source of truth for the TRT engine's ``action_horizon`` and
``batch_size``, read back from ``export_metadata.json``.
When a model is exported to ONNX/TensorRT, two numbers are baked into the
engine and recorded in ``export_metadata.json`` next to it:
* ``action_horizon`` — the predicted action-chunk length. It also determines
the engine's static ``sa_seq_len`` (``1 + action_horizon``).
* ``batch_size`` — baked as a *static* shape (the exporter registers only the
sequence dim in ``dynamic_axes``), so the engine only accepts that exact
batch at runtime.
The same two numbers are then re-stated independently elsewhere: in the
loaded model's config, in the ``--batch-size`` flag of the verify / benchmark
scripts, and in the ``--action-horizon`` open-loop stride of the standalone
inference script. When any copy drifts from the engine, the failure is silent
or cryptic — a foreign or stale ``.engine`` dropped into the bundle, or a
typo'd ``--batch-size``, surfaces only as a generic ``Invalid input shape``
raised deep inside the engine's ``forward()``, naming neither the engine's
baked value nor the user's flag.
The helpers here read the baked values back from ``export_metadata.json`` and
validate each re-stated copy against them up-front, with error messages that
name both sides. If the metadata file is missing (older bundles) the checks
degrade to a warning rather than failing.
"""
from __future__ import annotations
import json
import logging
import os
from typing import Any
logger = logging.getLogger(__name__)
_METADATA_FILENAME = "export_metadata.json"
# Bumped when export_metadata.json changes incompatibly (a key the build/runtime
# readers depend on is renamed, removed, or repurposed). The exporter stamps
# this; the build reader rejects a bundle whose version it does not recognize.
EXPORT_METADATA_SCHEMA_VERSION = 1
# Keys the build reader (build_tensorrt_engine.build_full_pipeline) needs to size
# the TRT shape profiles, plus the values the runtime contract re-checks. A bundle
# missing any of these cannot be built without guessing a sequence/patch shape.
# ``schema_version`` is intentionally not here: the version gate in
# validate_export_metadata handles its absence before the missing-keys check runs.
REQUIRED_EXPORT_METADATA_KEYS = (
"sa_seq_len",
"vl_seq_len",
"llm_seq_len",
"num_patches",
"num_merged_patches",
"num_vis_tokens",
"action_horizon",
"batch_size",
"precision",
)
def validate_export_metadata(
metadata: dict[str, Any],
*,
source: str = "export metadata",
engine_path: str = "",
) -> None:
"""Raise unless ``metadata`` is a current-schema, build-ready bundle.
Checks ``schema_version`` equals :data:`EXPORT_METADATA_SCHEMA_VERSION` and
every :data:`REQUIRED_EXPORT_METADATA_KEYS` entry is present, so a stale
bundle or a renamed/dropped field fails here — naming the cause — instead of
silently defaulting to a wrong sequence/patch hint deep in the TRT build. The
message states the problem; the caller decides the remedy.
"""
where = f" at {engine_path}" if engine_path else ""
version = metadata.get("schema_version")
if version != EXPORT_METADATA_SCHEMA_VERSION:
raise ValueError(
f"{source}: {_METADATA_FILENAME}{where} has schema_version={version!r}, but "
f"this build expects {EXPORT_METADATA_SCHEMA_VERSION}"
)
missing = [k for k in REQUIRED_EXPORT_METADATA_KEYS if k not in metadata]
if missing:
raise ValueError(
f"{source}: {_METADATA_FILENAME}{where} is missing required key(s) {missing}"
)
def _candidate_metadata_paths(engine_path: str) -> list[str]:
"""Locations to look for ``export_metadata.json`` given an engine path.
``engine_path`` may be an engine directory or a single ``.engine`` file
(dit_only mode). The metadata is written by ``export_onnx_n1d7`` into the
ONNX output dir and copied next to the engines by ``build_trt_pipeline``,
so we also check a sibling ``onnx/`` dir for un-copied legacy layouts.
"""
base = engine_path
if os.path.isfile(engine_path) or engine_path.endswith(".engine"):
base = os.path.dirname(engine_path)
candidates = [
os.path.join(base, _METADATA_FILENAME),
os.path.join(os.path.dirname(base.rstrip("/")), "onnx", _METADATA_FILENAME),
]
return candidates
def load_export_metadata(engine_path: str) -> dict[str, Any] | None:
"""Return the export metadata for an engine bundle, or ``None`` if absent.
A missing *or unreadable* (corrupt JSON / IO error) metadata file returns
``None`` so callers can degrade to a warning uniformly rather than crashing
on a malformed file.
"""
for path in _candidate_metadata_paths(engine_path):
if os.path.exists(path):
try:
with open(path) as f:
return json.load(f)
except (json.JSONDecodeError, OSError) as e:
logger.warning(
"Failed to read export metadata %s (%s); treating as absent.",
path,
e,
)
return None
return None
def _policy_action_horizon(policy: Any) -> int | None:
"""Best-effort read of the loaded policy's action horizon."""
action_head = getattr(getattr(policy, "model", None), "action_head", None)
if action_head is None:
return None
cfg = getattr(action_head, "config", None)
if cfg is not None and getattr(cfg, "action_horizon", None) is not None:
return int(cfg.action_horizon)
if getattr(action_head, "action_horizon", None) is not None:
return int(action_head.action_horizon)
return None
def assert_engine_matches_policy(
policy: Any,
engine_path: str,
*,
source: str = "setup_tensorrt_engines",
) -> dict[str, Any] | None:
"""Validate that an engine bundle was built for the loaded policy.
Compares the engine's recorded ``action_horizon`` (and the derived
``sa_seq_len == 1 + action_horizon``) against the loaded policy's action
head. A mismatch — e.g. a foreign or stale ``.engine`` dropped into the
bundle — raises here, naming both values, instead of surfacing as a
generic ``Invalid input shape`` deep inside ``Engine.forward()``.
When ``export_metadata.json`` is absent the contract cannot be checked; we
log a warning and return ``None`` rather than failing (older bundles).
"""
metadata = load_export_metadata(engine_path)
if metadata is None:
logger.warning(
"%s: no %s found next to %s; cannot validate that the engine's "
"action_horizon / batch_size match the loaded policy. A "
"mismatched engine will fail later as a cryptic 'Invalid input "
"shape' inside Engine.forward().",
source,
_METADATA_FILENAME,
engine_path,
)
return None
engine_ah = metadata.get("action_horizon")
engine_sa = metadata.get("sa_seq_len")
if engine_ah is not None and engine_sa is not None and engine_sa != engine_ah + 1:
raise ValueError(
f"{source}: corrupt {_METADATA_FILENAME} for {engine_path}: "
f"sa_seq_len={engine_sa} but action_horizon={engine_ah} "
f"(expected sa_seq_len == 1 + action_horizon == {engine_ah + 1})."
)
policy_ah = _policy_action_horizon(policy)
if engine_ah is not None and policy_ah is not None and engine_ah != policy_ah:
sa_note = f" (baked into sa_seq_len={engine_sa})" if engine_sa is not None else ""
raise ValueError(
f"{source}: TRT engine bundle at {engine_path} was built for "
f"action_horizon={engine_ah}{sa_note}, but the loaded policy has "
f"action_horizon={policy_ah}. The engine and policy disagree on "
"chunk size; re-export/rebuild the engines for this model, or load "
"the model the engines were built from."
)
return metadata
def assert_engine_bundle_present(
engine_path: str,
required_files,
*,
mode: str = "n17_full_pipeline",
source: str = "setup_tensorrt_engines",
) -> None:
"""Fail fast, with build instructions, when a TRT engine bundle is missing.
``setup_tensorrt_engines`` swaps in several ``.engine`` files that have no
PyTorch fallback (the action head's state/action encoders, the DiT, and the
action decoder). If the ``--trt-engine-path`` directory does not exist, or
exists but is missing one of those files, the loader would otherwise die with
a bare ``FileNotFoundError`` deep inside ``Engine.load`` — giving the user no
hint that they simply have not built the engines yet. Raise an actionable
error here instead, naming the missing directory / files and the build step.
"""
build_hint = (
"Build the engines first, e.g.:\n"
" python scripts/deployment/build_trt_pipeline.py \\\n"
" --model-path <model> --dataset-path <dataset> \\\n"
" --embodiment-tag <TAG> --output-dir ./gr00t_trt_deployment\n"
"then pass --trt-engine-path ./gr00t_trt_deployment/engines "
"(see scripts/deployment/ for the full deployment guide)."
)
if not os.path.isdir(engine_path):
raise FileNotFoundError(
f"{source}: inference-mode '{mode}' needs a TensorRT engine "
f"directory, but none exists at {engine_path!r}.\n{build_hint}"
)
missing = [f for f in required_files if not os.path.exists(os.path.join(engine_path, f))]
if missing:
raise FileNotFoundError(
f"{source}: inference-mode '{mode}' requires these TensorRT engine "
f"file(s) in {engine_path!r}, which are missing: "
f"{', '.join(sorted(missing))}.\n{build_hint}"
)
def resolve_batch_size(
engine_path: str,
requested: int | None = None,
*,
source: str = "TRT runtime",
) -> int:
"""Resolve the runtime batch size against the engine's build-time batch.
``export_onnx_n1d7`` bakes the batch dim as a static shape (only
``seq_len`` is in ``dynamic_axes``), so the engine only accepts the exact
batch it was built at. This reads that value from ``export_metadata.json``
and validates the requested batch against it:
- ``requested is None`` -> return the engine's build batch.
- ``requested != build batch`` -> raise, naming both (a typo'd
``--batch-size`` otherwise fails as a cryptic ``Invalid input shape``).
"""
metadata = load_export_metadata(engine_path)
built = metadata.get("batch_size") if metadata else None
if requested is None:
if built is None:
return 1
return int(built)
if built is not None and int(requested) != int(built):
raise ValueError(
f"{source}: requested batch_size={requested} but the TRT engine at "
f"{engine_path} was built (statically) for batch_size={built}. "
"The export pipeline does not register the batch dim in "
"dynamic_axes, so the engine only accepts its build batch. Pass "
f"--batch-size {built}, or rebuild the engines at batch_size="
f"{requested}."
)
return int(requested)
def assert_grid_thw_matches(
baked_grid: Any,
runtime_grid_thw: Any,
*,
source: str = "ViT TRT forward",
) -> None:
"""Validate a runtime ``image_grid_thw`` against the ViT engine's baked grid.
The ViT export pre-computes position/rotary embeddings for the captured
``grid_thw`` and freezes them as buffers; the engine's only input is
``pixel_values``. Those buffers depend on each view's ``[t, h, w]``
*layout*, not on how many views are present: batching tiles the same
per-view grid, so the runtime view count scales with batch size while the
layouts stay fixed (and a real total-shape mismatch is already rejected by
the static ``pixel_values`` shape). So we require every runtime view's
layout to be one the engine baked; a view with an unbaked layout (e.g. H/W
swapped, different resolution or temporal span) would get the wrong
embeddings and is rejected, while a different view *count* (batch) is fine.
``baked_grid`` is ``None`` for engine bundles built before ``vit_grid_thw``
was recorded; skip the check then (degrade like a missing metadata file).
"""
if baked_grid is None:
return
rt = runtime_grid_thw
if hasattr(rt, "detach"):
rt = rt.detach().cpu()
if hasattr(rt, "tolist"):
rt = rt.tolist()
rt_rows = [tuple(int(x) for x in row) for row in rt]
baked_layouts = {tuple(int(x) for x in row) for row in baked_grid}
unbaked = sorted({row for row in rt_rows if row not in baked_layouts})
if unbaked:
raise ValueError(
f"{source}: ViT TRT engine baked position/rotary buffers for "
f"image_grid_thw layouts {sorted(baked_layouts)}, but this observation "
f"has view layout(s) {[list(r) for r in unbaked]} that were never baked. "
f"The engine ignores runtime grid_thw (pixel_values is its only input) "
f"and would silently produce wrong vision features. Re-export/rebuild "
f"the ViT engine for this image configuration, or run with a baked layout."
)
def assert_exec_horizon_within_model(
*,
exec_horizon: int,
model_action_horizon: int,
source: str = "inference",
) -> None:
"""Validate an open-loop execution stride against the model's chunk size.
``standalone_inference_script --execution-horizon`` is the number of actions
consumed per predicted chunk; it must not exceed the model's
``action_horizon`` (the predicted chunk length), otherwise indexing the
chunk by ``range(exec_horizon)`` runs past the end.
"""
if not (1 <= exec_horizon <= model_action_horizon):
raise ValueError(
f"{source}: --execution-horizon={exec_horizon} must satisfy "
f"1 <= execution_horizon <= model action_horizon={model_action_horizon} "
"(= the predicted chunk length). A larger stride indexes past the "
"predicted action chunk."
)
__all__ = [
"EXPORT_METADATA_SCHEMA_VERSION",
"REQUIRED_EXPORT_METADATA_KEYS",
"validate_export_metadata",
"load_export_metadata",
"assert_engine_matches_policy",
"assert_engine_bundle_present",
"resolve_batch_size",
"assert_grid_thw_matches",
"assert_exec_horizon_within_model",
]
|