Skip to content

[MD-TRT] Llama and Qwen export TP examples with KV cache - #4784

Open
apbose wants to merge 11 commits into
mainfrom
abose/llama_qwen_tp_export_kvcache_refresh
Open

apbose wants to merge 11 commits into
mainfrom
abose/llama_qwen_tp_export_kvcache_refresh

Conversation

@apbose

@apbose apbose commented Oct 6, 2026 •

Copy link
Copy Markdown
Collaborator

Description

tools/llm/tensor_parallel_llm_export.py provides tensor-parallel export, save and load for both Llama and Qwen with static_v1 / static_v2 KV 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 default atol=rtol=0.02 and the measured Qwen2.5-0.5B-Instruct / static_v2 / TP=2 / batch=1 budget of atol=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.

  • 5 serialization regressions passed, including the standalone scalar-bounds regression.
  • Small Llama and Qwen3 FP32, static_v2: fresh export, accuracy check and separate-process load all passed (6/6 stages).
  • Existing small Qwen3 FP16 engines passed the default accuracy check.
  • The README accuracy snippet passed 128 steps on pretrained Qwen2.5-0.5B-Instruct engines in both FP32 and the documented FP16 configuration. Each matched 128/128 next-token choices against its corresponding full eager reference on each rank; maximum absolute logit differences were 0.0001354 (FP32) and 0.0546875 (FP16).
  • Black, Ruff and git diff --check pass for all changed Python files.

Pretrained Qwen was also freshly exported and reloaded with static_v1 in both FP16 and FP32. Its accuracy was replayed on the same 904 contexts as the earlier static_v2 experiment (six prompts x 64 steps, four x 128, plus eight synthetic-token steps), using identical input histories:

Cache Precision Maximum absolute logit difference Matching next-token choices Passing logit checks
static_v1 FP16 0.0957031 902/904 903/904
static_v1 FP32 0.0006673 904/904 904/904
static_v2 FP16 0.0805664 902/904 904/904
static_v2 FP32 0.0006111 904/904 904/904

Logit checks use atol=0.08, rtol=0.02 for FP16 and atol=0.001, rtol=0.0001 for FP32. One static_v1 FP16 logit exceeded the tolerance at the eighth synthetic-input step (required atol approximately 0.08412 at rtol=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_export still fails at the zero-input TRT partition assertion. Reproduced with both the rebased exporter and main 8b0cda5f64's exact exporter in the source-validation environment. The separate local fix is intentionally not included in this PR update.
  • The 904-context replay follows fixed reference token histories; it is not a free-running model-quality benchmark. FP32 performance was not measured.
  • No validation of TRT-RTX, multi-node execution, TP above two, pretrained Llama accuracy, or longer contexts beyond the reported runs.
  • Benchmark argument handling, TP benchmark input synchronization, batch-size profiles, uncached EOS stopping and aggregate throughput accounting remain separate review work; the tests above exercise non-benchmark, batch=1 paths.

@github-actions github-actions Bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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)

@apbose apbose added this to the v2.15.0 milestone Oct 6, 2026
@github-actions github-actions Bot added component: tests Issues re: Tests component: core Issues re: The core compiler component: api [Python] Issues re: Python API component: dynamo Issues relating to the `torch.compile` or `torch._dynamo.export` paths labels Oct 7, 2026

@github-actions github-actions Bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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)

apbose added 7 commits October 8, 2026 07:30
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
apbose force-pushed the abose/llama_qwen_tp_export_kvcache_refresh branch from 46481e2 to 85d33c6 Compare October 8, 2026 07:38
@github-actions github-actions Bot added the documentation Improvements or additions to documentation label Oct 8, 2026
@apbose
apbose force-pushed the abose/llama_qwen_tp_export_kvcache_refresh branch from 738286b to a1fffae Compare October 8, 2026 19:13
@apbose
apbose requested a review from narendasan October 9, 2026 20:43

@narendasan narendasan left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

LGTM

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

cla signed component: api [Python] Issues re: Python API component: core Issues re: The core compiler component: dynamo Issues relating to the `torch.compile` or `torch._dynamo.export` paths component: tests Issues re: Tests documentation Improvements or additions to documentation

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants