-
Notifications
You must be signed in to change notification settings - Fork 2.7k
[docs] Add VLA full hidden_states export example #17420
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
base: main
Are you sure you want to change the base?
Changes from all commits
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,81 @@ | ||
| # Exporting Full Hidden States from TensorRT-LLM for VLA Models | ||
|
|
||
| ## Background | ||
|
|
||
| VLA (Vision-Language-Action) models such as Orion, OpenVLA, and RT-2 use the | ||
| LLM's hidden states as input to downstream task heads (e.g. planning, trajectory | ||
| prediction) rather than for token generation: | ||
|
|
||
| - Standard LLM: input -> LLM -> logits -> token sampling -> text | ||
| - VLA model: images+text -> LLM -> hidden_states -> planning head -> trajectory | ||
|
|
||
| The downstream head needs the hidden state at a **specific token position** | ||
| (e.g. a "waypoint" token), which requires the complete hidden_states tensor, | ||
| not just the last token's logits. | ||
|
|
||
| See [issue #4414](https://github.com/NVIDIA/TensorRT-LLM/issues/4414) for | ||
| community demand (open since May 2025). | ||
|
|
||
| ## Why Existing APIs Fall Short | ||
|
|
||
| | API | Limitation | | ||
| |-----|-----------| | ||
| | `gather_last_token_logits` | Compresses to last token only | | ||
| | `additional_model_outputs` (v1.1+) | Requires model `forward` to return `hidden_states` in its output dict; the standard `DecoderModelForCausalLM.forward` only returns logits | | ||
| | `SaveHiddenStatesDecodingConfig` | Offline only (EAGLE3 training), saves to disk | | ||
|
|
||
| ## Solutions | ||
|
|
||
| ### Solution A: TRT Backend (v0.7-v0.21) | ||
|
|
||
| For TRT network-based builds, insert `mark_output` before | ||
| `gather_last_token_logits` in `modeling_utils.py`: | ||
|
|
||
| ```python | ||
| if self.config.mapping.is_last_pp_rank(): | ||
| # Export full hidden_states before compression | ||
| hidden_states.mark_output('full_hidden_states', self.config.dtype) | ||
| hidden_states = gather_last_token_logits(...) | ||
| lm_logits = self.lm_head(hidden_states) | ||
| ``` | ||
|
|
||
| **Note on tensor shape**: `mark_output` exposes the tensor as-is without | ||
| reshaping. With `remove_input_padding` enabled, the shape may be packed | ||
| `[num_tokens, hidden_dim]` rather than `[batch, seq_len, hidden_dim]`. | ||
|
|
||
| Reading at inference: | ||
|
|
||
| ```python | ||
| full_hs = model.session.debug_buffer["full_hidden_states"] | ||
| # Shape is [batch, seq_len, hidden_dim] or [num_tokens, hidden_dim] (packed) | ||
| ego_feature = full_hs[0, waypoint_idx, :] # or full_hs[waypoint_idx, :] | ||
| ``` | ||
|
|
||
| ### Solution B: PyTorch Backend (v1.x) | ||
|
|
||
| In v1.x, `DecoderModelForCausalLM.forward` returns logits only. To expose | ||
| hidden_states via `additional_model_outputs`, modify `forward` to return a dict | ||
| when requested. The key insight: `self.model()` returns the full tensor | ||
| before `LogitsProcessor` compresses it. See `patches/modeling_utils_v1x.patch` | ||
| for the approach, and refer to `handle_additional_outputs.py` for the | ||
| framework's dict-return contract. | ||
|
|
||
| **Note**: On the PyTorch backend with `remove_input_padding` (default), the | ||
| shape is packed `[num_tokens, hidden_dim]`. | ||
|
|
||
| ## Verification | ||
|
|
||
| Solution A was tested on Orion VLA (ICCV 2025) with TRT-LLM v0.13.0: | ||
|
|
||
| - Engine output: `[1, 599, 4096]` | ||
| - Hidden_states CosSim vs PyTorch: 0.9994 (INT8) | ||
| - End-to-end plan_L2_1s: 0.686 (PyTorch: 0.690) | ||
|
|
||
| ## Files | ||
|
|
||
| | File | Description | | ||
| |------|-------------| | ||
| | `patches/modeling_utils_v0x.patch` | Solution A (v0.7-v0.21, verified) | | ||
| | `patches/modeling_utils_v1x.patch` | Solution B (v1.x, conceptual) | | ||
| | `inference_python.py` | Python inference example | | ||
| | `tests/test_hidden_states.py` | Tests | |
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,28 @@ | ||
| # SPDX-FileCopyrightText: Copyright (c) 2024-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. | ||
| # SPDX-License-Identifier: Apache-2.0 | ||
| """ | ||
| Build a TRT-LLM engine with full_hidden_states output (v0.x TRT backend). | ||
|
|
||
| NOTE: This script applies to the legacy TRT backend (v0.7-v0.21) which uses | ||
| `trtllm-build`. On v1.x (main branch), the build command has changed to | ||
| `trtllm serve` / `trtllm bench`. For v1.x, refer to Solution B in the README | ||
| (modify model forward to return hidden_states dict). | ||
|
Comment on lines
+6
to
+9
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. 🎯 Functional Correctness | 🟠 Major | ⚡ Quick win 🧩 Analysis chain🏁 Script executed: #!/bin/bash
set -eu
printf '%s\n' '--- target file ---'
cat -n examples/vla_hidden_states_export/export_hidden_states.py | sed -n '1,30p'
printf '%s\n' '--- repository references ---'
rg -n --glob '!node_modules' --glob '!build' --glob '!dist' 'trtllm-build|trtllm-serve|trtllm-bench|trtllm serve|trtllm bench|v0\.7|v0\.21|v1\.x' .
printf '%s\n' '--- referenced upstream README ---'
curl -L --fail --silent --show-error --max-time 20 \
'https://raw.githubusercontent.com/NVIDIA/TensorRT-LLM/main/examples/models/core/qwen/README.md' |
rg -n -C 3 'trtllm-build|trtllm-serve|trtllm-bench|trtllm serve|trtllm bench|version|release' || trueRepository: NVIDIA/TensorRT-LLM Length of output: 50378 🏁 Script executed: #!/bin/bash
set -eu
printf '%s\n' '--- example directory ---'
git ls-files examples/vla_hidden_states_export
printf '%s\n' '--- example README and related references ---'
if test -f examples/vla_hidden_states_export/README.md; then
cat -n examples/vla_hidden_states_export/README.md | sed -n '1,240p'
fi
rg -n -C 4 'Solution B|hidden_states|v0\.7|v0\.21|v1\.x|trtllm-build|trtllm-serve|trtllm-bench' \
examples/vla_hidden_states_export README.md docs examples/models/core/qwen/README.md
printf '%s\n' '--- focused local Qwen README ---'
cat -n examples/models/core/qwen/README.md | sed -n '1,190p'
printf '%s\n' '--- upstream Qwen README focused extract ---'
curl -L --fail --silent --show-error --max-time 20 \
'https://raw.githubusercontent.com/NVIDIA/TensorRT-LLM/main/examples/models/core/qwen/README.md' |
sed -n '1,190p'
printf '%s\n' '--- package/version and release references ---'
rg -n --glob '!.git/**' '__version__|version\s*=|v0\.[0-9]+|v1\.[0-9]+' setup.py pyproject.toml tensorrt_llm 2>/dev/null |
head -200 || true
printf '%s\n' '--- recent tags/releases visible locally ---'
git tag --sort=-version:refname | head -30Repository: NVIDIA/TensorRT-LLM Length of output: 50376 🏁 Script executed: #!/bin/bash
set -eu
printf '%s\n' '--- tracked files in example ---'
git ls-files examples/vla_hidden_states_export
printf '%s\n' '--- example README ---'
if test -f examples/vla_hidden_states_export/README.md; then
cat -n examples/vla_hidden_states_export/README.md | sed -n '1,220p'
else
printf '%s\n' 'README.md not present'
fi
printf '%s\n' '--- exact local Qwen command references ---'
rg -n -C 5 'trtllm-build|trtllm-serve|trtllm-bench|Solution B|hidden_states' \
examples/models/core/qwen/README.md examples/vla_hidden_states_export/export_hidden_states.py
printf '%s\n' '--- exact legacy benchmark build documentation ---'
rg -n -C 4 'trtllm-bench.*build|trtllm-build' \
docs/source/legacy/performance/perf-benchmarking.md \
docs/source/legacy/performance/performance-tuning-guide/benchmarking-default-performance.md |
head -120
printf '%s\n' '--- release/version metadata ---'
git tag --sort=-version:refname | head -40
printf '%s\n' '--- repository version declarations ---'
rg -n --glob 'setup.py' --glob 'pyproject.toml' --glob 'version.py' \
'version\s*=|__version__' . | head -80 || true
printf '%s\n' '--- upstream Qwen README command lines ---'
curl -L --fail --silent --show-error --max-time 20 \
'https://raw.githubusercontent.com/NVIDIA/TensorRT-LLM/main/examples/models/core/qwen/README.md' |
rg -n -C 5 'trtllm-build|trtllm-serve|trtllm-bench|Solution B|hidden_states'Repository: NVIDIA/TensorRT-LLM Length of output: 39018 Correct the v1.x CLI description and pin the legacy release. For v1.x, the PyTorch backend loads Hugging Face checkpoints directly and has no engine-build step. Use 🤖 Prompt for AI Agents |
||
|
|
||
| Prerequisites for v0.x: | ||
| 1. Apply patches/modeling_utils_v0x.patch to tensorrt_llm/models/modeling_utils.py | ||
| 2. This adds `hidden_states.mark_output('full_hidden_states', ...)` before | ||
| gather_last_token_logits, preserving the full 3D tensor. | ||
|
|
||
| Verify the engine has 'full_hidden_states' as an output after build: | ||
|
|
||
| python -c " | ||
| import tensorrt as trt | ||
| runtime = trt.Runtime(trt.Logger()) | ||
| with open('rank0.engine', 'rb') as f: | ||
| engine = runtime.deserialize_cuda_engine(f.read()) | ||
| outputs = [engine.get_tensor_name(i) for i in range(engine.num_io_tensors) | ||
| if engine.get_tensor_mode(engine.get_tensor_name(i)) == trt.TensorIOMode.OUTPUT] | ||
| print('Outputs:', outputs) | ||
| assert 'full_hidden_states' in outputs, 'mark_output patch not applied!' | ||
| " | ||
| """ | ||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,140 @@ | ||
| # SPDX-FileCopyrightText: Copyright (c) 2024-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. | ||
| # SPDX-License-Identifier: Apache-2.0 | ||
| """ | ||
| Minimal example: read full_hidden_states from a TRT-LLM engine at inference time. | ||
|
|
||
| This shows the core usage pattern for VLA models: | ||
| 1. Run LLM engine inference | ||
| 2. Extract the full hidden_states from engine output | ||
| 3. Select the hidden state at a specific token position (e.g., waypoint token) | ||
| 4. Feed to downstream planning head | ||
| """ | ||
| import argparse | ||
| import numpy as np | ||
| import tensorrt as trt | ||
| import pycuda.driver as cuda | ||
| import pycuda.autoinit | ||
|
|
||
|
|
||
| class HiddenStatesEngine: | ||
| """Minimal TRT engine wrapper that reads full_hidden_states output.""" | ||
|
|
||
| def __init__(self, engine_path: str) -> None: | ||
| logger = trt.Logger(trt.Logger.WARNING) | ||
| with open(engine_path, "rb") as f: | ||
| engine = trt.Runtime(logger).deserialize_cuda_engine(f.read()) | ||
|
|
||
| self.engine = engine | ||
| self.context = engine.create_execution_context() | ||
| self.stream = cuda.Stream() | ||
|
|
||
| self.input_names = [] | ||
| self.output_names = [] | ||
| for i in range(engine.num_io_tensors): | ||
| name = engine.get_tensor_name(i) | ||
| if engine.get_tensor_mode(name) == trt.TensorIOMode.INPUT: | ||
| self.input_names.append(name) | ||
| else: | ||
| self.output_names.append(name) | ||
|
|
||
| if "full_hidden_states" not in self.output_names: | ||
| raise ValueError( | ||
| f"full_hidden_states not in engine outputs: {self.output_names}. " | ||
| "Apply the mark_output patch first." | ||
| ) | ||
| print(f"Engine loaded: {len(self.input_names)} inputs, {len(self.output_names)} outputs") | ||
|
|
||
| def infer_and_extract_hidden_states( | ||
| self, | ||
| input_ids: np.ndarray, | ||
| waypoint_idx: int, | ||
| ) -> tuple[np.ndarray, np.ndarray]: | ||
| """ | ||
| Run inference and extract hidden_states at the waypoint token position. | ||
|
|
||
| Args: | ||
| input_ids: Input token IDs (numpy int32 array, shape [batch, seq_len]) | ||
| waypoint_idx: Position in the sequence to extract hidden state from. | ||
|
|
||
| Returns: | ||
| full_hidden_states: Complete hidden_states tensor. | ||
| ego_feature: Hidden state at waypoint position [hidden_dim]. | ||
| """ | ||
| buffers = {} | ||
|
|
||
| for name in self.input_names: | ||
| if name == "input_ids": | ||
| data = np.ascontiguousarray(input_ids.astype(np.int32)) | ||
| self.context.set_input_shape(name, data.shape) | ||
| else: | ||
| shape = tuple(max(1, s) for s in self.engine.get_tensor_shape(name)) | ||
| data = np.zeros(shape, dtype=np.float32) | ||
|
|
||
| d = cuda.mem_alloc(data.nbytes) | ||
| cuda.memcpy_htod(d, data) | ||
| self.context.set_tensor_address(name, int(d)) | ||
| buffers[name] = d | ||
|
|
||
| for name in self.output_names: | ||
| engine_shape = tuple(self.engine.get_tensor_shape(name)) | ||
| shape = tuple(max(1, s) for s in engine_shape) | ||
| n_elements = int(np.prod(shape)) | ||
| d = cuda.mem_alloc(n_elements * 2) # fp16 | ||
| self.context.set_tensor_address(name, int(d)) | ||
| buffers[name] = d | ||
|
|
||
| self.context.execute_async_v3(self.stream.handle) | ||
| self.stream.synchronize() | ||
|
|
||
| hs_shape = tuple( | ||
| max(1, s) for s in self.engine.get_tensor_shape("full_hidden_states") | ||
| ) | ||
| full_hs = np.empty(int(np.prod(hs_shape)), dtype=np.float16) | ||
| cuda.memcpy_dtoh(full_hs, buffers["full_hidden_states"]) | ||
| full_hs = full_hs.reshape(hs_shape) | ||
|
|
||
| if full_hs.ndim == 3: # [batch, seq_len, hidden_dim] | ||
| wp_idx = min(waypoint_idx, full_hs.shape[1] - 1) | ||
| ego_feature = full_hs[0, wp_idx, :].copy() | ||
| elif full_hs.ndim == 2: # [num_tokens, hidden_dim] (packed, remove_input_padding) | ||
| wp_idx = min(waypoint_idx, full_hs.shape[0] - 1) | ||
| ego_feature = full_hs[wp_idx, :].copy() | ||
| else: | ||
| ego_feature = full_hs[-1, :].copy() | ||
|
|
||
| for d in buffers.values(): | ||
| d.free() | ||
|
|
||
| return full_hs, ego_feature | ||
|
|
||
|
|
||
| def main() -> None: | ||
| parser = argparse.ArgumentParser( | ||
| description="Read full_hidden_states from TRT-LLM engine" | ||
| ) | ||
| parser.add_argument( | ||
| "--engine_path", required=True, help="Path to rank0.engine" | ||
| ) | ||
| parser.add_argument( | ||
| "--seq_len", type=int, default=599, | ||
| help="Sequence length for dummy input (default: 599)", | ||
| ) | ||
| parser.add_argument( | ||
| "--waypoint_idx", type=int, default=598, | ||
| help="Token position to extract (default: last token)", | ||
| ) | ||
| args = parser.parse_args() | ||
|
|
||
| engine = HiddenStatesEngine(args.engine_path) | ||
| input_ids = np.ones((1, args.seq_len), dtype=np.int32) | ||
| full_hs, ego_feature = engine.infer_and_extract_hidden_states( | ||
| input_ids, args.waypoint_idx | ||
| ) | ||
|
|
||
| print(f"\nfull_hidden_states shape: {full_hs.shape}") | ||
| print(f"ego_feature (idx={args.waypoint_idx}): {ego_feature.shape}") | ||
| print(f"ego_feature[:6]: {ego_feature[:6]}") | ||
|
|
||
|
|
||
| if __name__ == "__main__": | ||
| main() |
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,26 @@ | ||
| From 0000000000000000000000000000000000000000 Mon Sep 17 00:00:00 2001 | ||
| From: Hou Song <hous.ailab@gmail.com> | ||
| Date: Thu, 7 Aug 2026 22:00:00 +0800 | ||
| Subject: [PATCH] Export full 3D hidden_states via mark_output | ||
|
|
||
| Insert mark_output('full_hidden_states') before gather_last_token_logits | ||
| in DecoderModelForCausalLM.forward(), so the complete tensor is exported. | ||
|
|
||
| Related issue: https://github.com/NVIDIA/TensorRT-LLM/issues/4414 | ||
|
|
||
| NOTE: Line numbers below are from v0.13.0. For other v0.x versions, search | ||
| for "gather_last_token_logits" in DecoderModelForCausalLM.forward() and | ||
| insert the mark_output line above it. The surrounding code is identical | ||
| across v0.7.0 through v0.21.0. | ||
|
|
||
| diff --git a/tensorrt_llm/models/modeling_utils.py b/tensorrt_llm/models/modeling_utils.py | ||
| --- a/tensorrt_llm/models/modeling_utils.py | ||
| +++ b/tensorrt_llm/models/modeling_utils.py | ||
| @@ -776,6 +776,10 @@ | ||
| if self.config.mapping.is_last_pp_rank(): | ||
| + # [VLA] Export full hidden_states before gather_last_token_logits compresses it. | ||
| + # See: https://github.com/NVIDIA/TensorRT-LLM/issues/4414 | ||
| + hidden_states.mark_output('full_hidden_states', self.config.dtype) | ||
| + | ||
| hidden_states = gather_last_token_logits( | ||
| hidden_states, last_token_ids, |
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,23 @@ | ||
| From 0000000000000000000000000000000000000000 Mon Sep 17 00:00:00 2001 | ||
| From: Hou Song <hous.ailab@gmail.com> | ||
| Date: Thu, 7 Aug 2026 22:00:00 +0800 | ||
| Subject: [PATCH] Return hidden_states as additional output (v1.x, conceptual) | ||
|
|
||
| In v1.x, DecoderModelForCausalLM.forward() calls self.model() which returns | ||
| the full hidden_states tensor, then passes it to LogitsProcessor which takes | ||
| only the last token. | ||
|
|
||
| The framework's additional_model_outputs contract (handle_additional_outputs.py) | ||
| expects the model forward to return a dict containing the requested keys. | ||
| The model_engine.py invocation passes additional_model_outputs to the forward | ||
| call. To expose hidden_states, modify forward to include it in the returned | ||
| dict before LogitsProcessor compresses it. | ||
|
|
||
| NOTE: On the PyTorch backend with remove_input_padding (default), the shape | ||
| is packed [num_tokens, hidden_dim], not [batch, seq_len, hidden_dim]. | ||
|
|
||
| This is a conceptual patch based on source code analysis of v1.x main branch. | ||
| Not yet runtime tested. The exact integration depends on how model_engine.py | ||
| plumbs the additional_model_outputs parameter into the model forward call. | ||
|
|
||
| Related issue: https://github.com/NVIDIA/TensorRT-LLM/issues/4414 |
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,85 @@ | ||
| # SPDX-FileCopyrightText: Copyright (c) 2024-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. | ||
| # SPDX-License-Identifier: Apache-2.0 | ||
| """ | ||
| Verify full_hidden_states output from a TRT-LLM engine. | ||
|
|
||
| Usage: | ||
| python tests/test_hidden_states.py --engine_path /path/to/rank0.engine | ||
| """ | ||
| import argparse | ||
| import numpy as np | ||
|
|
||
|
|
||
| def test_engine_has_full_hidden_states(engine_path: str) -> None: | ||
| """Verify the engine exports full_hidden_states tensor.""" | ||
| import tensorrt as trt | ||
|
|
||
| logger = trt.Logger(trt.Logger.WARNING) | ||
| with open(engine_path, "rb") as f: | ||
| engine = trt.Runtime(logger).deserialize_cuda_engine(f.read()) | ||
|
|
||
| output_names = [ | ||
| engine.get_tensor_name(i) | ||
| for i in range(engine.num_io_tensors) | ||
| if engine.get_tensor_mode(engine.get_tensor_name(i)) | ||
| == trt.TensorIOMode.OUTPUT | ||
| ] | ||
|
|
||
| if "full_hidden_states" not in output_names: | ||
| raise AssertionError( | ||
| f"full_hidden_states not in outputs: {output_names}. " | ||
| "Apply the mark_output patch first." | ||
| ) | ||
| print(f"PASS: full_hidden_states found in {output_names}") | ||
|
|
||
|
|
||
| def test_hidden_states_is_3d(engine_path: str) -> None: | ||
| """Verify full_hidden_states has rank >= 2 (3D or packed 2D).""" | ||
| import tensorrt as trt | ||
|
|
||
| logger = trt.Logger(trt.Logger.WARNING) | ||
| with open(engine_path, "rb") as f: | ||
| engine = trt.Runtime(logger).deserialize_cuda_engine(f.read()) | ||
|
|
||
| shape = tuple(engine.get_tensor_shape("full_hidden_states")) | ||
| if len(shape) < 2: | ||
| raise AssertionError( | ||
| f"full_hidden_states should be at least 2D, got shape {shape}." | ||
| ) | ||
| print(f"PASS: full_hidden_states shape {shape} (rank {len(shape)})") | ||
|
|
||
|
|
||
| def test_token_extraction() -> None: | ||
| """Verify token extraction from a simulated full hidden_states tensor.""" | ||
| seq_len = 599 | ||
| hidden_dim = 4096 | ||
| full_hs = np.random.randn(1, seq_len, hidden_dim).astype(np.float16) | ||
|
|
||
| wp_idx = 42 | ||
| ego_feature = full_hs[0, wp_idx, :].copy() | ||
|
|
||
| assert ego_feature.shape == (hidden_dim,) | ||
| assert np.array_equal(ego_feature, full_hs[0, wp_idx, :]) | ||
| assert not np.array_equal(ego_feature, full_hs[0, wp_idx + 1, :]) | ||
| print(f"PASS: token extraction at idx={wp_idx}") | ||
|
|
||
|
|
||
| def main() -> None: | ||
| import sys | ||
| parser = argparse.ArgumentParser() | ||
| parser.add_argument("--engine_path", required=True) | ||
| args = parser.parse_args() | ||
|
|
||
| test_engine_has_full_hidden_states(args.engine_path) | ||
| test_hidden_states_is_3d(args.engine_path) | ||
| test_token_extraction() | ||
| print("\nAll tests passed!") | ||
|
|
||
|
|
||
| if __name__ == "__main__": | ||
| import sys | ||
| if "--engine_path" in sys.argv: | ||
| main() | ||
| else: | ||
| test_token_extraction() | ||
| print("\nRun with --engine_path to test engine output") |
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
New files need the NVIDIA copyright header (
AGENTS.md, "CRITICAL" section). Applies toinference_python.pyandtests/test_hidden_states.pytoo.