Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
31 changes: 31 additions & 0 deletions tensorrt_llm/_torch/pyexecutor/cuda_graph_runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -370,6 +370,37 @@ def maybe_get_cuda_graph(
graph_spec_metadata = None
return graph_attn_metadata, graph_spec_metadata, key

def clear_capture_only_spec_state(self) -> int:
"""Clear capture-scoped state from every cached graph SpecMetadata.

``create_cuda_graph_metadata`` shallow-copies the live SpecMetadata, so a
copy made while ``_run_capture_pass(force_non_greedy=True)`` is active
inherits ``_force_non_greedy_for_capture=True``. That copy is cached here
and reseated as the live spec_metadata on every later replay of its graph,
while the capture pass clears the flag on the base object only. Without
this cleanup the copies keep the flag forever and
``_scan_one_model_sampling`` rewrites EVERY serving request's sampling
params to the synthetic capture values (temperature 0.7 / top_k 50 /
top_p 0.9), silently ignoring what the client asked for.

The flag must NOT be cleared at copy time instead: it is load-bearing
*during* capture. It is what makes the pass-2 populate scan non-greedy on
parameter-less warmup requests, so that the advanced-sampling branch (not
the argmax fast path, and with the top-k/top-p kernels present) is the one
recorded into the graph. Clearing it here -- after the pass has captured
every graph -- keeps capture correct and serving clean.

Returns the number of cached metadata objects cleared.
"""
cleared = 0
for stored in self.graph_metadata.values():
spec_metadata = stored.get("spec_metadata")
if spec_metadata is not None and getattr(
spec_metadata, "_force_non_greedy_for_capture", False):
spec_metadata._force_non_greedy_for_capture = False
cleared += 1
return cleared

def needs_capture(self, key: KeyType):
return self._capture_allowed and key not in self.graph_outputs

Expand Down
25 changes: 25 additions & 0 deletions tensorrt_llm/_torch/pyexecutor/model_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -1987,6 +1987,21 @@ def _run_capture_pass(force_non_greedy: bool, label: str) -> None:
finally:
if force_non_greedy and spec_metadata is not None:
spec_metadata._force_non_greedy_for_capture = False
# The base object is not the only holder of the flag: every
# graph captured during this pass cached its own SHALLOW COPY
# of spec_metadata (create_cuda_graph_metadata -> copy.copy),
# which inherited the flag. Those copies are reseated as the
# live spec_metadata on every later replay, so leaving the
# flag set there makes _scan_one_model_sampling overwrite
# every serving request's sampling params with the synthetic
# capture values (0.7 / 50 / 0.9). Clear them here -- after
# the pass has finished capturing, so the flag was still in
# effect for every capture that needed it.
cleared = self.cuda_graph_runner.clear_capture_only_spec_state(
)
logger.info(
f"Cleared capture-only sampling override from {cleared} "
"cached CUDA graph spec metadata object(s).")
Comment thread
xwang233 marked this conversation as resolved.

# Pass 1: greedy fast-path (dummy requests carry no sampling params,
# so is_all_greedy_sample is naturally True).
Expand Down Expand Up @@ -4983,6 +4998,16 @@ def previous_seq_slots_device():
num_accepted_draft_tokens)]
if isinstance(spec_metadata, Eagle3SpecMetadata):
spec_metadata.request_accepted_path = request_accepted_path
# The capture-only sampling override must never be live outside CUDA
# graph warmup: it replaces every request's sampling params with
# synthetic capture values. It leaked here once already (inherited by
# the cached graph metadata shallow copies), so assert rather than
# trust the teardown.
assert self.is_warmup or not getattr(
spec_metadata, '_force_non_greedy_for_capture', False
), ("capture-only sampling override (_force_non_greedy_for_capture) "
"is set outside CUDA graph warmup; serving requests would be "
"silently decoded with the synthetic capture sampling params")
# No-op for non 1-model
spec_metadata.populate_sampling_params_for_one_model(
scheduled_requests.all_requests())
Expand Down
1 change: 1 addition & 0 deletions tests/integration/test_lists/test-db/l0_h100.yml
Original file line number Diff line number Diff line change
Expand Up @@ -50,6 +50,7 @@ l0_h100:
- unittest/_torch/sampler -k "not test_speculative_d2h_parity_real_predictor"
- unittest/_torch/speculative/test_eagle3.py
- unittest/_torch/speculative/test_rejection_buffers_guard.py
- unittest/_torch/speculative/test_capture_override_leak.py
- unittest/_torch/speculative/test_sa_hybrid_state_promotion.py
- unittest/_torch/speculative/hw_agnostic
- unittest/_torch/thop/parallel
Expand Down
143 changes: 143 additions & 0 deletions tests/unittest/_torch/speculative/test_capture_override_leak.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,143 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
"""CPU unit tests for the lifetime of the capture-only sampling override.

The advanced-sampling CUDA graph capture pass sets
``_force_non_greedy_for_capture=True`` on the live ``SpecMetadata`` so that
parameter-less warmup requests scan as non-greedy and the advanced-sampling
branch is the one recorded into the graph.

``create_cuda_graph_metadata`` shallow-copies the live metadata, so every graph
captured during that pass caches a copy that inherited the flag, and those
copies are reseated as the live spec_metadata on every later replay. Clearing
the flag on the base object alone therefore leaves it set forever on the copies,
and ``_scan_one_model_sampling`` then rewrites EVERY serving request's sampling
params to the synthetic capture values.

These tests use a real (base) ``SpecMetadata`` -- its ``__post_init__`` is a
no-op and none of the fields exercised here are tensors -- plus an unbound call
of ``CUDAGraphRunner.clear_capture_only_spec_state`` on a stand-in holding only
``graph_metadata``, mirroring test_group_all_greedy_sync.py. No GPU, no runner
construction, and no model forward is needed.
"""

import types

import torch

from tensorrt_llm._torch.pyexecutor.cuda_graph_runner import CUDAGraphRunner
from tensorrt_llm._torch.pyexecutor.llm_request import LlmRequestState
from tensorrt_llm._torch.speculative.interface import SpecMetadata

# The synthetic params the capture override substitutes for the request's own,
# and the sentinel that means "top-k disabled" (see _scan_one_model_sampling).
CAPTURE_TEMPERATURE, CAPTURE_TOP_K, CAPTURE_TOP_P = 0.7, 50, 0.9
DISABLE_TOPK_VAL = torch.iinfo(torch.int32).max


def _base_meta():
"""A live (non-graph) SpecMetadata, as the model engine holds it."""
return SpecMetadata(max_num_requests=8, max_draft_len=1, max_total_draft_tokens=1)


def _graph_copy(meta, batch_size=8):
"""The shallow copy maybe_get_cuda_graph caches for one captured graph."""
graph_meta = meta.create_cuda_graph_metadata(batch_size)
assert graph_meta is not meta
return graph_meta


def _clear(graph_metadata):
"""CUDAGraphRunner.clear_capture_only_spec_state, called unbound."""
return CUDAGraphRunner.clear_capture_only_spec_state(
types.SimpleNamespace(graph_metadata=graph_metadata)
)


def _request(temperature=None, top_k=None, top_p=None, slot=0):
return types.SimpleNamespace(
sampling_config=types.SimpleNamespace(
temperature=[temperature] if temperature is not None else None,
top_k=[top_k] if top_k is not None else None,
top_p=[top_p] if top_p is not None else None,
),
state=LlmRequestState.GENERATION_IN_PROGRESS,
py_seq_slot=slot,
)


def _scan(meta, requests):
normalized, _ = SpecMetadata._scan_one_model_sampling(meta, requests)
# Drop the trailing num_tokens; only the sampling params matter here.
return [entry[:3] for entry in normalized]


def test_graph_copy_inherits_flag_and_base_teardown_does_not_reach_it():
# The mechanism the bug rests on: copy.copy carries the flag over, and the
# copies are independent objects, so clearing the base misses them.
meta = _base_meta()
meta._force_non_greedy_for_capture = True
copies = [_graph_copy(meta, bs) for bs in (1, 2, 4)]
assert all(copy._force_non_greedy_for_capture for copy in copies)

meta._force_non_greedy_for_capture = False
assert all(copy._force_non_greedy_for_capture for copy in copies)


def test_clear_capture_only_spec_state_clears_every_cached_copy():
meta = _base_meta()
meta._force_non_greedy_for_capture = True
advanced = [_graph_copy(meta, bs) for bs in (1, 2, 4)]
# Graphs captured by the greedy pass never had the flag, and non-spec
# graphs cache no spec_metadata at all; both must be left alone.
meta._force_non_greedy_for_capture = False
greedy = _graph_copy(meta, 8)

graph_metadata = {("greedy", 8): {"spec_metadata": greedy}}
graph_metadata[("no_spec", 1)] = {"spec_metadata": None}
for i, copy in enumerate(advanced):
graph_metadata[("advanced", i)] = {"spec_metadata": copy}

assert _clear(graph_metadata) == len(advanced)
assert not any(copy._force_non_greedy_for_capture for copy in advanced)
assert greedy._force_non_greedy_for_capture is False
# Idempotent: a second teardown finds nothing left to clear.
assert _clear(graph_metadata) == 0


def test_serving_scan_honors_client_params_after_capture_teardown():
# End-to-end property of the fix, and the case that fails without it: with
# only the base-object teardown this scan returns (0.7, 50, 0.9).
meta = _base_meta()
meta._force_non_greedy_for_capture = True
graph_meta = _graph_copy(meta)

meta._force_non_greedy_for_capture = False # base-object teardown
_clear({("advanced", 8): {"spec_metadata": graph_meta}}) # the fix

# Replay reseats the cached copy as the live spec_metadata, so the serving
# scan runs on it, not on the base object.
assert _scan(graph_meta, [_request(temperature=1.0, top_p=1.0)]) == [
(1.0, DISABLE_TOPK_VAL, 1.0)
]
assert graph_meta.is_all_greedy_sample is False


def test_override_stays_live_while_the_flag_is_set():
# Anti-regression for the rejected "clear at copy time" fix: the flag is
# load-bearing *during* capture. Cleared any earlier, the pass-2 populate
# would scan these parameter-less warmup requests as greedy and bake the
# argmax fast path -- with no top-k/top-p kernels -- into the graph keyed
# as the advanced-sampling variant.
meta = _base_meta()
meta._force_non_greedy_for_capture = True
graph_meta = _graph_copy(meta)

warmup_requests = [_request(slot=None), _request(slot=None)]
assert _scan(graph_meta, warmup_requests) == [
(CAPTURE_TEMPERATURE, CAPTURE_TOP_K, CAPTURE_TOP_P)
] * len(warmup_requests)
assert graph_meta.is_all_greedy_sample is False
assert not graph_meta.skip_temperature
assert not graph_meta.skip_top_k
assert not graph_meta.skip_top_p
Loading