Repository navigation
Conversation
There was a problem hiding this comment.
There are some changes that do not conform to Python style guidelines:
--- /home/runner/work/TensorRT/TensorRT/tools/llm/utils.py 2026-10-06 06:20:29.334643+00:00
+++ /home/runner/work/TensorRT/TensorRT/tools/llm/utils.py 2026-10-06 06:21:07.957711+00:00
@@ -210,13 +210,13 @@
"""
Greedy decoding of the model with static KV cache.
"""
start_idx = 0
end_idx = input_seq.shape[1]
- position_ids = torch.arange(
- input_seq.shape[1], device=input_seq.device
- ).unsqueeze(0)
+ position_ids = torch.arange(input_seq.shape[1], device=input_seq.device).unsqueeze(
+ 0
+ )
output_seq = input_seq.clone()
# TODO: Confirm this: When end_idx = max_output_seq_length-1, number of tokens generated = OSL
num_tokens_generated = 0
kv_cache = get_zeroed_static_cache_inputs(model, device=input_seq.device)
while end_idx < max_output_seq_length:
@@ -328,13 +328,13 @@
time reflects GPU completion, not just kernel launches.
"""
start_idx = 0
end_idx = input_seq.shape[1]
prefill_tokens = end_idx
- position_ids = torch.arange(
- input_seq.shape[1], device=input_seq.device
- ).unsqueeze(0)
+ position_ids = torch.arange(input_seq.shape[1], device=input_seq.device).unsqueeze(
+ 0
+ )
kv_cache = get_zeroed_static_cache_inputs(model, device=input_seq.device)
torch.cuda.synchronize()
prefill_start = timeit.default_timer()
input_signature = (input_seq, position_ids, *kv_cache, start_idx, end_idx)There was a problem hiding this comment.
There are some changes that do not conform to Python style guidelines:
--- /home/runner/work/TensorRT/TensorRT/tools/llm/utils.py 2026-10-07 05:48:38.438493+00:00
+++ /home/runner/work/TensorRT/TensorRT/tools/llm/utils.py 2026-10-07 05:49:18.004551+00:00
@@ -210,13 +210,13 @@
"""
Greedy decoding of the model with static KV cache.
"""
start_idx = 0
end_idx = input_seq.shape[1]
- position_ids = torch.arange(
- input_seq.shape[1], device=input_seq.device
- ).unsqueeze(0)
+ position_ids = torch.arange(input_seq.shape[1], device=input_seq.device).unsqueeze(
+ 0
+ )
output_seq = input_seq.clone()
# TODO: Confirm this: When end_idx = max_output_seq_length-1, number of tokens generated = OSL
num_tokens_generated = 0
kv_cache = get_zeroed_static_cache_inputs(model, device=input_seq.device)
while end_idx < max_output_seq_length:
@@ -328,13 +328,13 @@
time reflects GPU completion, not just kernel launches.
"""
start_idx = 0
end_idx = input_seq.shape[1]
prefill_tokens = end_idx
- position_ids = torch.arange(
- input_seq.shape[1], device=input_seq.device
- ).unsqueeze(0)
+ position_ids = torch.arange(input_seq.shape[1], device=input_seq.device).unsqueeze(
+ 0
+ )
kv_cache = get_zeroed_static_cache_inputs(model, device=input_seq.device)
torch.cuda.synchronize()
prefill_start = timeit.default_timer()
input_signature = (input_seq, position_ids, *kv_cache, start_idx, end_idx)Keep cache tensors and position IDs on the input device in generation and split-timing helpers. This fixes tensor-parallel ranks on nonzero GPUs allocating their cache on cuda:0.
Collect range constraints from standalone SymInt placeholders as well as tensor dimensions in the legacy exporter. Cache start/end indices otherwise lose their bounds during save(retrace=False), making the saved program fail to reload. Add a regression covering tensor and scalar bounds, save/load, and execution with different scalar values.
Normalize cache-enabled TP models to plain FX CodeGen before saving. Cache lowering adds positional KV/index inputs and flat outputs, while the retained pytree metadata still requires a position_ids keyword and describes the old output structure. Let the legacy exporter rebuild the saved input and output structure from the transformed graph so a loaded model accepts the cache decoder call.
apbose
force-pushed
the
abose/llama_qwen_tp_export_kvcache_refresh
branch
from
October 8, 2026 07:38
46481e2 to
85d33c6
Compare
apbose
force-pushed
the
abose/llama_qwen_tp_export_kvcache_refresh
branch
from
October 8, 2026 19:13
738286b to
a1fffae
Compare
This branch has not been deployed
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Description
tools/llm/tensor_parallel_llm_export.pyprovides tensor-parallel export, save and load for both Llama and Qwen withstatic_v1/static_v2KV caches. Cache tensors and position IDs stay on the input device, including ranks on nonzero GPUs. Supersedes #4287.Cached models could compile and run in memory but fail after saving. The serialization fixes preserve bounds for standalone scalar inputs such as
start_idx/end_idx, and derive the saved positional KV/index inputs and flat outputs from the transformed graph rather than retaining the original model's calling structure. A regression covers scalar bounds, save/load and execution with different scalar values. Previously saved broken cache artifacts must be re-exported.The example now accepts
--precision fp16|fp32(FP16 remains the default). FP32 converts weights and buffers before tracing, disables autocast, and uses the same precision during export, compilation and inference. TensorRT compilation disables TF32. Changing a saved engine's precision requires re-exporting into a separate directory.The README's Accuracy check section contains a runnable code snippet comparing saved cached engines against an unsharded eager reference with identical token histories. It reports logit differences and next-token agreement separately. Export/load do not take tolerance arguments.
The snippet uses FP32 tolerances
atol=0.001, rtol=0.0001. Its FP16 instructions give the defaultatol=rtol=0.02and the measured Qwen2.5-0.5B-Instruct /static_v2/ TP=2 / batch=1 budget ofatol=0.08, rtol=0.02.The README's Precision findings section presents the experiment settings, TRT-versus-full-eager results, and the two prompt/generated-token comparisons. A separate small helper cleanup removes an unused import, narrows a bare exception handler, and fixes formatting.
Validation
Rebased onto main
8b0cda5f64611f3a060963f42dd330446731652f. Post-rebase checks used the full rebased Python package with the existing validation container's compiled Torch-TensorRT libraries; this was not a fresh C++/wheel build of latest main. Environment: standard TensorRT 11.3.0.99, PyTorch 2.15.0.dev20261005+cu130, Transformers 5.14.1, and 2 x B300 per distributed case. TensorRT matches main's pin; the prepared build uses CUDA 13.2 rather than main's named CUDA 13.4 toolkit.static_v2: fresh export, accuracy check and separate-process load all passed (6/6 stages).0.0001354(FP32) and0.0546875(FP16).git diff --checkpass for all changed Python files.Pretrained Qwen was also freshly exported and reloaded with
static_v1in both FP16 and FP32. Its accuracy was replayed on the same 904 contexts as the earlierstatic_v2experiment (six prompts x 64 steps, four x 128, plus eight synthetic-token steps), using identical input histories:static_v1static_v1static_v2static_v2Logit checks use
atol=0.08, rtol=0.02for FP16 andatol=0.001, rtol=0.0001for FP32. Onestatic_v1FP16 logit exceeded the tolerance at the eighth synthetic-input step (requiredatolapproximately 0.08412 atrtol=0.02); its next-token choice still matched. The threshold was not relaxed.Both cache variants had the same two FP16 token mismatches: the FP16 eager reference's top scores tied, while FP16 TRT's selected tokens agreed with both FP32 implementations. Token agreement does not mean identical logits or establish factual answer quality.
Earlier validation covered both static cache variants for small Llama/Qwen3 models: 12/12 export/load/accuracy stages passed and 40/40 rank/step comparisons met the default FP16 tolerance.
Remaining issues and coverage limits
test_arange_exportstill fails at the zero-input TRT partition assertion. Reproduced with both the rebased exporter and main8b0cda5f64's exact exporter in the source-validation environment. The separate local fix is intentionally not included in this PR update.