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",
]