From 23640662c3591d95db1628213a63f25304987af4 Mon Sep 17 00:00:00 2001 From: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> Date: Thu, 27 Aug 2026 07:32:31 -0700 Subject: [PATCH 01/23] [WiP] Allow disabling torch.compile for generation-only batches Signed-off-by: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> --- tensorrt_llm/_torch/pyexecutor/model_engine.py | 15 ++++++++++++++- tensorrt_llm/llmapi/llm_args.py | 7 +++++++ .../unittest/llmapi/test_torch_compile_config.py | 12 ++++++++++++ 3 files changed, 33 insertions(+), 1 deletion(-) create mode 100644 tests/unittest/llmapi/test_torch_compile_config.py diff --git a/tensorrt_llm/_torch/pyexecutor/model_engine.py b/tensorrt_llm/_torch/pyexecutor/model_engine.py index 2429877c9cb4..fa128dbce86c 100644 --- a/tensorrt_llm/_torch/pyexecutor/model_engine.py +++ b/tensorrt_llm/_torch/pyexecutor/model_engine.py @@ -638,6 +638,8 @@ def __init__( torch_compile_enabled = bool(self.torch_compile_config is not None) torch_compile_fullgraph = self.torch_compile_config.enable_fullgraph if self.torch_compile_config is not None else TorchCompileConfig.model_fields[ 'enable_fullgraph'].default + torch_compile_generation = self.torch_compile_config.compile_generation if self.torch_compile_config is not None else TorchCompileConfig.model_fields[ + 'compile_generation'].default torch_compile_inductor_enabled = self.torch_compile_config.enable_inductor if self.torch_compile_config is not None else TorchCompileConfig.model_fields[ 'enable_inductor'].default torch_compile_piecewise_cuda_graph = (self.prefill_cuda_graph_backend == @@ -703,17 +705,28 @@ def __init__( fullgraph=torch_compile_fullgraph) self._torch_compile_prefill_only = ( self._torch_compile_piecewise_cuda_graph - and not self.model.use_fx_for_pcg_fallback) + and (not torch_compile_generation + or not self.model.use_fx_for_pcg_fallback)) self.model.model = ( _PrefillCompiledModel(eager_model, compiled_model) if self._torch_compile_prefill_only else compiled_model) elif callable(apply_llm_torch_compile): + if not torch_compile_generation: + raise ValueError( + "TorchCompileConfig.compile_generation=False is " + "only supported for DecoderModelForCausalLM models." + ) # TODO: Move this contract to MultimodalModelMixin once # multimodal models consistently expose their LLM compile # scope through the mixin. apply_llm_torch_compile(backend=self._torch_compile_backend, fullgraph=torch_compile_fullgraph) else: + if not torch_compile_generation: + raise ValueError( + "TorchCompileConfig.compile_generation=False is " + "only supported for DecoderModelForCausalLM models." + ) self.model = torch.compile( self.model, backend=self._torch_compile_backend, diff --git a/tensorrt_llm/llmapi/llm_args.py b/tensorrt_llm/llmapi/llm_args.py index 0e68df417237..30d88c0d3d7a 100644 --- a/tensorrt_llm/llmapi/llm_args.py +++ b/tensorrt_llm/llmapi/llm_args.py @@ -5763,6 +5763,13 @@ class TorchCompileConfig(StrictBaseModel): default=True, description="Enable full graph compilation in torch.compile.") + compile_generation: bool = Field( + default=True, + description="Apply torch.compile to generation-only batches. When " + "disabled, context and mixed batches remain compiled while " + "generation-only batches use eager execution.", + status="prototype") + enable_inductor: bool = Field( default=False, description="Enable inductor backend in torch.compile.") diff --git a/tests/unittest/llmapi/test_torch_compile_config.py b/tests/unittest/llmapi/test_torch_compile_config.py new file mode 100644 index 000000000000..aa2b525a6b64 --- /dev/null +++ b/tests/unittest/llmapi/test_torch_compile_config.py @@ -0,0 +1,12 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +from tensorrt_llm.llmapi.llm_args import TorchCompileConfig + + +def test_compile_generation_defaults_to_enabled(): + assert TorchCompileConfig().compile_generation + + +def test_compile_generation_can_be_disabled(): + assert not TorchCompileConfig(compile_generation=False).compile_generation From ee82d07fe0add809c5bb580d40a5ade85144aaff Mon Sep 17 00:00:00 2001 From: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> Date: Fri, 28 Aug 2026 01:56:49 -0700 Subject: [PATCH 02/23] Refactor compile_generation flag to cover warmup and autotuning passes Signed-off-by: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> --- .../_torch/pyexecutor/model_engine.py | 43 ++++++++++++++----- tensorrt_llm/llmapi/llm_args.py | 11 ++--- .../llmapi/test_torch_compile_config.py | 9 ++-- 3 files changed, 43 insertions(+), 20 deletions(-) diff --git a/tensorrt_llm/_torch/pyexecutor/model_engine.py b/tensorrt_llm/_torch/pyexecutor/model_engine.py index fa128dbce86c..6b14d9394f87 100644 --- a/tensorrt_llm/_torch/pyexecutor/model_engine.py +++ b/tensorrt_llm/_torch/pyexecutor/model_engine.py @@ -638,8 +638,10 @@ def __init__( torch_compile_enabled = bool(self.torch_compile_config is not None) torch_compile_fullgraph = self.torch_compile_config.enable_fullgraph if self.torch_compile_config is not None else TorchCompileConfig.model_fields[ 'enable_fullgraph'].default - torch_compile_generation = self.torch_compile_config.compile_generation if self.torch_compile_config is not None else TorchCompileConfig.model_fields[ - 'compile_generation'].default + compile_only_context_and_mixed_graphs = ( + self.torch_compile_config.compile_only_context_and_mixed_graphs + if self.torch_compile_config is not None else TorchCompileConfig. + model_fields['compile_only_context_and_mixed_graphs'].default) torch_compile_inductor_enabled = self.torch_compile_config.enable_inductor if self.torch_compile_config is not None else TorchCompileConfig.model_fields[ 'enable_inductor'].default torch_compile_piecewise_cuda_graph = (self.prefill_cuda_graph_backend == @@ -650,6 +652,10 @@ def __init__( 'max_num_streams'].default self._torch_compile_enabled = torch_compile_enabled + self._compile_only_context_and_mixed_graphs = ( + compile_only_context_and_mixed_graphs) + torch_compile_bypass_state = {"active": False} + self._torch_compile_bypass_state = torch_compile_bypass_state self._torch_compile_piecewise_cuda_graph = torch_compile_piecewise_cuda_graph self._torch_compile_prefill_only = False @@ -705,15 +711,16 @@ def __init__( fullgraph=torch_compile_fullgraph) self._torch_compile_prefill_only = ( self._torch_compile_piecewise_cuda_graph - and (not torch_compile_generation + and (compile_only_context_and_mixed_graphs or not self.model.use_fx_for_pcg_fallback)) self.model.model = ( _PrefillCompiledModel(eager_model, compiled_model) if self._torch_compile_prefill_only else compiled_model) elif callable(apply_llm_torch_compile): - if not torch_compile_generation: + if compile_only_context_and_mixed_graphs: raise ValueError( - "TorchCompileConfig.compile_generation=False is " + "TorchCompileConfig." + "compile_only_context_and_mixed_graphs=True is " "only supported for DecoderModelForCausalLM models." ) # TODO: Move this contract to MultimodalModelMixin once @@ -722,9 +729,10 @@ def __init__( apply_llm_torch_compile(backend=self._torch_compile_backend, fullgraph=torch_compile_fullgraph) else: - if not torch_compile_generation: + if compile_only_context_and_mixed_graphs: raise ValueError( - "TorchCompileConfig.compile_generation=False is " + "TorchCompileConfig." + "compile_only_context_and_mixed_graphs=True is " "only supported for DecoderModelForCausalLM models." ) self.model = torch.compile( @@ -1295,6 +1303,19 @@ def _pad_batch_seed_mrope_delta_cache( for request in mrope_seed_requests: request.py_mrope_delta_cache_slot = request.py_seq_slot + @contextmanager + def _without_torch_compile(self): + if not self._compile_only_context_and_mixed_graphs: + yield + return + state = self._torch_compile_bypass_state + previous = state["active"] + state["active"] = True + try: + yield + finally: + state["active"] = previous + @staticmethod def warmup_with_kv_cache_cleanup(method): """ @@ -1531,7 +1552,7 @@ def _warmup_scheduled(self, resource_manager: ResourceManager, with self._warmup_timer.phase( "attention_jit", metrics=self._metrics, - metric_name="attention_warmup_seconds"): + metric_name="attention_warmup_seconds"), self._without_torch_compile(): self._run_attention_warmup(resource_manager, can_run_general_warmup) @@ -1561,7 +1582,7 @@ def _warmup_scheduled(self, resource_manager: ResourceManager, with self._warmup_timer.phase( "autotuner", metrics=self._metrics, - metric_name="autotuner_warmup_seconds"): + metric_name="autotuner_warmup_seconds"), self._without_torch_compile(): self._run_autotuner_warmup(resource_manager) log_mem_snapshot("warmup/after_autotuner") # Pre-JIT Mamba SSD multi-seq + HAS_INITSTATES=True Triton kernels @@ -1572,7 +1593,7 @@ def _warmup_scheduled(self, resource_manager: ResourceManager, with self._warmup_timer.phase( "mamba_hybrid", metrics=self._metrics, - metric_name="mamba_hybrid_warmup_seconds"): + metric_name="mamba_hybrid_warmup_seconds"), self._without_torch_compile(): self._run_mamba_hybrid_warmup(resource_manager) log_mem_snapshot("warmup/after_mamba_hybrid") # Release the autotuner's exploration-mode intermediates. The @@ -1625,7 +1646,7 @@ def _warmup_scheduled(self, resource_manager: ResourceManager, with self._warmup_timer.phase( "memory_pool_prepop", metrics=self._metrics, - metric_name="memory_pool_prepopulation_seconds"): + metric_name="memory_pool_prepopulation_seconds"), self._without_torch_compile(): warmup_requests_configs = self._get_max_shape_warmup_requests( resource_manager) self._general_warmup(resource_manager, warmup_requests_configs) diff --git a/tensorrt_llm/llmapi/llm_args.py b/tensorrt_llm/llmapi/llm_args.py index 30d88c0d3d7a..9c4de3583e27 100644 --- a/tensorrt_llm/llmapi/llm_args.py +++ b/tensorrt_llm/llmapi/llm_args.py @@ -5763,11 +5763,12 @@ class TorchCompileConfig(StrictBaseModel): default=True, description="Enable full graph compilation in torch.compile.") - compile_generation: bool = Field( - default=True, - description="Apply torch.compile to generation-only batches. When " - "disabled, context and mixed batches remain compiled while " - "generation-only batches use eager execution.", + compile_only_context_and_mixed_graphs: bool = Field( + default=False, + description="Compile only context and mixed-batch execution graphs. " + "When enabled, generation-only and auxiliary kernel warmup forwards " + "use eager execution, while dedicated context/mixed compilation and " + "piecewise CUDA graph warmup and capture remain compiled.", status="prototype") enable_inductor: bool = Field( diff --git a/tests/unittest/llmapi/test_torch_compile_config.py b/tests/unittest/llmapi/test_torch_compile_config.py index aa2b525a6b64..d93e58fcd168 100644 --- a/tests/unittest/llmapi/test_torch_compile_config.py +++ b/tests/unittest/llmapi/test_torch_compile_config.py @@ -4,9 +4,10 @@ from tensorrt_llm.llmapi.llm_args import TorchCompileConfig -def test_compile_generation_defaults_to_enabled(): - assert TorchCompileConfig().compile_generation +def test_compile_only_context_and_mixed_graphs_defaults_to_disabled(): + assert not TorchCompileConfig().compile_only_context_and_mixed_graphs -def test_compile_generation_can_be_disabled(): - assert not TorchCompileConfig(compile_generation=False).compile_generation +def test_compile_only_context_and_mixed_graphs_can_be_enabled(): + config = TorchCompileConfig(compile_only_context_and_mixed_graphs=True) + assert config.compile_only_context_and_mixed_graphs From 82fe10a82cd00d48344401de0607d8959eab931f Mon Sep 17 00:00:00 2001 From: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> Date: Mon, 31 Aug 2026 05:39:44 -0700 Subject: [PATCH 03/23] Rename new flag to compile_only_piecewise_graphs Signed-off-by: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> --- .../_torch/pyexecutor/model_engine.py | 21 +++++++++---------- tensorrt_llm/llmapi/llm_args.py | 2 +- 2 files changed, 11 insertions(+), 12 deletions(-) diff --git a/tensorrt_llm/_torch/pyexecutor/model_engine.py b/tensorrt_llm/_torch/pyexecutor/model_engine.py index 6b14d9394f87..f9062ec69083 100644 --- a/tensorrt_llm/_torch/pyexecutor/model_engine.py +++ b/tensorrt_llm/_torch/pyexecutor/model_engine.py @@ -638,10 +638,10 @@ def __init__( torch_compile_enabled = bool(self.torch_compile_config is not None) torch_compile_fullgraph = self.torch_compile_config.enable_fullgraph if self.torch_compile_config is not None else TorchCompileConfig.model_fields[ 'enable_fullgraph'].default - compile_only_context_and_mixed_graphs = ( - self.torch_compile_config.compile_only_context_and_mixed_graphs + compile_only_piecewise_graphs = ( + self.torch_compile_config.compile_only_piecewise_graphs if self.torch_compile_config is not None else TorchCompileConfig. - model_fields['compile_only_context_and_mixed_graphs'].default) + model_fields['compile_only_piecewise_graphs'].default) torch_compile_inductor_enabled = self.torch_compile_config.enable_inductor if self.torch_compile_config is not None else TorchCompileConfig.model_fields[ 'enable_inductor'].default torch_compile_piecewise_cuda_graph = (self.prefill_cuda_graph_backend == @@ -652,8 +652,7 @@ def __init__( 'max_num_streams'].default self._torch_compile_enabled = torch_compile_enabled - self._compile_only_context_and_mixed_graphs = ( - compile_only_context_and_mixed_graphs) + self._compile_only_piecewise_graphs = compile_only_piecewise_graphs torch_compile_bypass_state = {"active": False} self._torch_compile_bypass_state = torch_compile_bypass_state self._torch_compile_piecewise_cuda_graph = torch_compile_piecewise_cuda_graph @@ -711,16 +710,16 @@ def __init__( fullgraph=torch_compile_fullgraph) self._torch_compile_prefill_only = ( self._torch_compile_piecewise_cuda_graph - and (compile_only_context_and_mixed_graphs + and (compile_only_piecewise_graphs or not self.model.use_fx_for_pcg_fallback)) self.model.model = ( _PrefillCompiledModel(eager_model, compiled_model) if self._torch_compile_prefill_only else compiled_model) elif callable(apply_llm_torch_compile): - if compile_only_context_and_mixed_graphs: + if compile_only_piecewise_graphs: raise ValueError( "TorchCompileConfig." - "compile_only_context_and_mixed_graphs=True is " + "compile_only_piecewise_graphs=True is " "only supported for DecoderModelForCausalLM models." ) # TODO: Move this contract to MultimodalModelMixin once @@ -729,10 +728,10 @@ def __init__( apply_llm_torch_compile(backend=self._torch_compile_backend, fullgraph=torch_compile_fullgraph) else: - if compile_only_context_and_mixed_graphs: + if compile_only_piecewise_graphs: raise ValueError( "TorchCompileConfig." - "compile_only_context_and_mixed_graphs=True is " + "compile_only_piecewise_graphs=True is " "only supported for DecoderModelForCausalLM models." ) self.model = torch.compile( @@ -1305,7 +1304,7 @@ def _pad_batch_seed_mrope_delta_cache( @contextmanager def _without_torch_compile(self): - if not self._compile_only_context_and_mixed_graphs: + if not self._compile_only_piecewise_graphs: yield return state = self._torch_compile_bypass_state diff --git a/tensorrt_llm/llmapi/llm_args.py b/tensorrt_llm/llmapi/llm_args.py index 9c4de3583e27..d4c2fbb87be4 100644 --- a/tensorrt_llm/llmapi/llm_args.py +++ b/tensorrt_llm/llmapi/llm_args.py @@ -5763,7 +5763,7 @@ class TorchCompileConfig(StrictBaseModel): default=True, description="Enable full graph compilation in torch.compile.") - compile_only_context_and_mixed_graphs: bool = Field( + compile_only_piecewise_graphs: bool = Field( default=False, description="Compile only context and mixed-batch execution graphs. " "When enabled, generation-only and auxiliary kernel warmup forwards " From 3cfea35b4403a60f3ea7f548edce5584f11fdb0b Mon Sep 17 00:00:00 2001 From: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> Date: Mon, 31 Aug 2026 05:54:56 -0700 Subject: [PATCH 04/23] Minor refactoring Signed-off-by: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> --- .../_torch/pyexecutor/model_engine.py | 20 ++++++++----------- 1 file changed, 8 insertions(+), 12 deletions(-) diff --git a/tensorrt_llm/_torch/pyexecutor/model_engine.py b/tensorrt_llm/_torch/pyexecutor/model_engine.py index f9062ec69083..b111e73e6ddf 100644 --- a/tensorrt_llm/_torch/pyexecutor/model_engine.py +++ b/tensorrt_llm/_torch/pyexecutor/model_engine.py @@ -702,6 +702,14 @@ def __init__( apply_llm_torch_compile = getattr(self.model, "apply_llm_torch_compile", None) + + if compile_only_piecewise_graphs and not isinstance( + self.model, DecoderModelForCausalLM): + raise ValueError( + "TorchCompileConfig." + "compile_only_piecewise_graphs=True is " + "only supported for DecoderModelForCausalLM models.") + if isinstance(self.model, DecoderModelForCausalLM): eager_model = self.model.model compiled_model = torch.compile( @@ -716,24 +724,12 @@ def __init__( _PrefillCompiledModel(eager_model, compiled_model) if self._torch_compile_prefill_only else compiled_model) elif callable(apply_llm_torch_compile): - if compile_only_piecewise_graphs: - raise ValueError( - "TorchCompileConfig." - "compile_only_piecewise_graphs=True is " - "only supported for DecoderModelForCausalLM models." - ) # TODO: Move this contract to MultimodalModelMixin once # multimodal models consistently expose their LLM compile # scope through the mixin. apply_llm_torch_compile(backend=self._torch_compile_backend, fullgraph=torch_compile_fullgraph) else: - if compile_only_piecewise_graphs: - raise ValueError( - "TorchCompileConfig." - "compile_only_piecewise_graphs=True is " - "only supported for DecoderModelForCausalLM models." - ) self.model = torch.compile( self.model, backend=self._torch_compile_backend, From 7ccba0d7ea2a0a1c5605486d54a6aa1df1a090cc Mon Sep 17 00:00:00 2001 From: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> Date: Mon, 31 Aug 2026 06:31:53 -0700 Subject: [PATCH 05/23] Encapsulate torch.compile bypass logic in helper class Signed-off-by: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> --- tensorrt_llm/_torch/compilation/utils.py | 38 ++++++++++++++- .../_torch/pyexecutor/model_engine.py | 15 ++---- .../test_phase_selective_forward.py | 48 +++++++++++++++++++ 3 files changed, 90 insertions(+), 11 deletions(-) create mode 100644 tests/unittest/_torch/compilation/test_phase_selective_forward.py diff --git a/tensorrt_llm/_torch/compilation/utils.py b/tensorrt_llm/_torch/compilation/utils.py index 2e98607b84b6..9daabc396e46 100644 --- a/tensorrt_llm/_torch/compilation/utils.py +++ b/tensorrt_llm/_torch/compilation/utils.py @@ -1,11 +1,47 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + import contextlib -from typing import Callable, List, Optional, Union +from collections.abc import Callable, Iterator +from typing import List, Optional, Union import torch from torch.fx import Node from torch.fx.experimental.symbolic_shapes import ShapeEnv from ..cuda_tile_utils import IS_CUDA_TILE_AVAILABLE +from ..utils import get_model_extra_attrs + + +class _PhaseSelectiveForward: + """Dispatch context/mixed batches to a compiled decoder.""" + + def __init__(self, eager_forward: Callable[..., object], + compiled_forward: Callable[..., object]) -> None: + self._eager_forward = eager_forward + self._compiled_forward = compiled_forward + self._bypass_active = False + + def __call__(self, *args, **kwargs) -> object: + attrs = get_model_extra_attrs() + attn_metadata_ref = attrs.get("attention_metadata") if attrs else None + attn_metadata = (attn_metadata_ref() + if attn_metadata_ref is not None else None) + # This is the finalized batch classification after any + # context-to-generation promotion. + if (self._bypass_active or + (attn_metadata is not None and attn_metadata.num_contexts == 0)): + return self._eager_forward(*args, **kwargs) + return self._compiled_forward(*args, **kwargs) + + @contextlib.contextmanager + def bypass(self) -> Iterator[None]: + previous = self._bypass_active + self._bypass_active = True + try: + yield + finally: + self._bypass_active = previous def get_symint_val(i: Union[torch.SymInt | int]): diff --git a/tensorrt_llm/_torch/pyexecutor/model_engine.py b/tensorrt_llm/_torch/pyexecutor/model_engine.py index b111e73e6ddf..511af829aac9 100644 --- a/tensorrt_llm/_torch/pyexecutor/model_engine.py +++ b/tensorrt_llm/_torch/pyexecutor/model_engine.py @@ -49,7 +49,8 @@ from ..autotuner import AutoTuner, autotune from ..compilation.backend import Backend from ..compilation.piecewise_optimizer import PiecewiseRunner -from ..compilation.utils import capture_piecewise_cuda_graph +from ..compilation.utils import (_PhaseSelectiveForward, + capture_piecewise_cuda_graph) from ..distributed import Distributed from ..distributed.communicator import init_pp_comm from ..memory_buffer_utils import clear_memory_buffers, with_shared_pool @@ -653,8 +654,7 @@ def __init__( self._torch_compile_enabled = torch_compile_enabled self._compile_only_piecewise_graphs = compile_only_piecewise_graphs - torch_compile_bypass_state = {"active": False} - self._torch_compile_bypass_state = torch_compile_bypass_state + self._phase_selective_forward: Optional[_PhaseSelectiveForward] = None self._torch_compile_piecewise_cuda_graph = torch_compile_piecewise_cuda_graph self._torch_compile_prefill_only = False @@ -1300,16 +1300,11 @@ def _pad_batch_seed_mrope_delta_cache( @contextmanager def _without_torch_compile(self): - if not self._compile_only_piecewise_graphs: + if self._phase_selective_forward is None: yield return - state = self._torch_compile_bypass_state - previous = state["active"] - state["active"] = True - try: + with self._phase_selective_forward.bypass(): yield - finally: - state["active"] = previous @staticmethod def warmup_with_kv_cache_cleanup(method): diff --git a/tests/unittest/_torch/compilation/test_phase_selective_forward.py b/tests/unittest/_torch/compilation/test_phase_selective_forward.py new file mode 100644 index 000000000000..d5f78b80ca4b --- /dev/null +++ b/tests/unittest/_torch/compilation/test_phase_selective_forward.py @@ -0,0 +1,48 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +from types import SimpleNamespace +from unittest.mock import Mock, patch + +from tensorrt_llm._torch.compilation.utils import _PhaseSelectiveForward + + +def _attention_attrs(num_contexts: int) -> dict[str, object]: + metadata = SimpleNamespace(num_contexts=num_contexts) + return {"attention_metadata": lambda: metadata} + + +def test_phase_selective_forward_dispatches_by_batch_phase(): + eager_forward = Mock(return_value="eager") + compiled_forward = Mock(return_value="compiled") + forward = _PhaseSelectiveForward(eager_forward, compiled_forward) + + with patch( + "tensorrt_llm._torch.compilation.utils.get_model_extra_attrs", + return_value=_attention_attrs(num_contexts=1), + ): + assert forward("context") == "compiled" + + with patch( + "tensorrt_llm._torch.compilation.utils.get_model_extra_attrs", + return_value=_attention_attrs(num_contexts=0), + ): + assert forward("generation") == "eager" + + compiled_forward.assert_called_once_with("context") + eager_forward.assert_called_once_with("generation") + + +def test_phase_selective_forward_bypass_is_restored(): + eager_forward = Mock(return_value="eager") + compiled_forward = Mock(return_value="compiled") + forward = _PhaseSelectiveForward(eager_forward, compiled_forward) + + with patch( + "tensorrt_llm._torch.compilation.utils.get_model_extra_attrs", + return_value=_attention_attrs(num_contexts=1), + ): + assert forward() == "compiled" + with forward.bypass(): + assert forward() == "eager" + assert forward() == "compiled" From c740f3735db4d95ff6f4c09123f193995bcbb5b8 Mon Sep 17 00:00:00 2001 From: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> Date: Mon, 31 Aug 2026 06:34:22 -0700 Subject: [PATCH 06/23] Drop trivial config tests Signed-off-by: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> --- tests/unittest/llmapi/test_torch_compile_config.py | 13 ------------- 1 file changed, 13 deletions(-) delete mode 100644 tests/unittest/llmapi/test_torch_compile_config.py diff --git a/tests/unittest/llmapi/test_torch_compile_config.py b/tests/unittest/llmapi/test_torch_compile_config.py deleted file mode 100644 index d93e58fcd168..000000000000 --- a/tests/unittest/llmapi/test_torch_compile_config.py +++ /dev/null @@ -1,13 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 - -from tensorrt_llm.llmapi.llm_args import TorchCompileConfig - - -def test_compile_only_context_and_mixed_graphs_defaults_to_disabled(): - assert not TorchCompileConfig().compile_only_context_and_mixed_graphs - - -def test_compile_only_context_and_mixed_graphs_can_be_enabled(): - config = TorchCompileConfig(compile_only_context_and_mixed_graphs=True) - assert config.compile_only_context_and_mixed_graphs From 70bd3bd277e3b8964173eaa96e4686c4ce7d26d3 Mon Sep 17 00:00:00 2001 From: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> Date: Mon, 31 Aug 2026 07:01:26 -0700 Subject: [PATCH 07/23] Add documentation Signed-off-by: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> --- tensorrt_llm/_torch/compilation/utils.py | 15 ++++++++++++++- tensorrt_llm/_torch/pyexecutor/model_engine.py | 5 +++++ 2 files changed, 19 insertions(+), 1 deletion(-) diff --git a/tensorrt_llm/_torch/compilation/utils.py b/tensorrt_llm/_torch/compilation/utils.py index 9daabc396e46..009e9d4fc07d 100644 --- a/tensorrt_llm/_torch/compilation/utils.py +++ b/tensorrt_llm/_torch/compilation/utils.py @@ -14,7 +14,12 @@ class _PhaseSelectiveForward: - """Dispatch context/mixed batches to a compiled decoder.""" + """ + This utility class is used to selectively bypass torch.compile + for operations that do not need it (attention warmup, auto-tuning), + as well as for generation-only forwards, where the ordinary CUDA + graph machinery is deemed sufficient. + """ def __init__(self, eager_forward: Callable[..., object], compiled_forward: Callable[..., object]) -> None: @@ -23,6 +28,11 @@ def __init__(self, eager_forward: Callable[..., object], self._bypass_active = False def __call__(self, *args, **kwargs) -> object: + """ + The wrapper will enforce the eager forward function if: + 1) The torch compile bypass is active (e.g., for auto-tuning), or + 2) The batch is generation-only (rely on ordinary CUDA graphs) + """ attrs = get_model_extra_attrs() attn_metadata_ref = attrs.get("attention_metadata") if attrs else None attn_metadata = (attn_metadata_ref() @@ -36,6 +46,9 @@ def __call__(self, *args, **kwargs) -> object: @contextlib.contextmanager def bypass(self) -> Iterator[None]: + """ + Disable torch.compile for all forwards under this ctxt manager. + """ previous = self._bypass_active self._bypass_active = True try: diff --git a/tensorrt_llm/_torch/pyexecutor/model_engine.py b/tensorrt_llm/_torch/pyexecutor/model_engine.py index 511af829aac9..bd532bd33102 100644 --- a/tensorrt_llm/_torch/pyexecutor/model_engine.py +++ b/tensorrt_llm/_torch/pyexecutor/model_engine.py @@ -1300,6 +1300,11 @@ def _pad_batch_seed_mrope_delta_cache( @contextmanager def _without_torch_compile(self): + """ + When compile_only_piecewise_graphs is enabled in TorchCompileConfig, + this ctxt manager bypasses torch.compile invocation anywhere outside + of piecewise CUDA graph capture and warmup. + """ if self._phase_selective_forward is None: yield return From 8ec255a663ad56ec87668410143e44719fd553bf Mon Sep 17 00:00:00 2001 From: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> Date: Tue, 1 Sep 2026 01:46:14 -0700 Subject: [PATCH 08/23] Add comments to justify missing torch.compile bypass Signed-off-by: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> --- tensorrt_llm/_torch/pyexecutor/model_engine.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/tensorrt_llm/_torch/pyexecutor/model_engine.py b/tensorrt_llm/_torch/pyexecutor/model_engine.py index bd532bd33102..0198afdd84e7 100644 --- a/tensorrt_llm/_torch/pyexecutor/model_engine.py +++ b/tensorrt_llm/_torch/pyexecutor/model_engine.py @@ -1560,6 +1560,8 @@ def _warmup_scheduled(self, resource_manager: ResourceManager, self._get_full_general_warmup_requests(resource_manager)) # Currently graph has not been captured, disable cuda graph for this warmup. with self.no_cuda_graph(): + # Do not bypass torch.compile in order to specialize the + # piecewise graphs before capture. self._general_warmup(resource_manager, warmup_requests_configs) # Release C++ MoE workspace buffers so the autotuner can @@ -1637,7 +1639,8 @@ def _warmup_scheduled(self, resource_manager: ResourceManager, log_mem_snapshot("warmup/after_cute_dsl_radix_topk") if can_run_general_warmup: # Pre-populate the memory pool with max-shape allocations to reduce - # fragmentation at runtime. + # fragmentation at runtime. If compile_only_piecewise_graphs is + # enabled, torch.compile can be safely bypassed here. with self._warmup_timer.phase( "memory_pool_prepop", metrics=self._metrics, From 6f2f62ff2b21e372a6c4a69d5e5580f1c9babfcf Mon Sep 17 00:00:00 2001 From: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> Date: Tue, 1 Sep 2026 13:37:52 +0000 Subject: [PATCH 09/23] Update LLM args manifest Signed-off-by: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> --- tensorrt_llm/usage/llm_args_golden_manifest.json | 7 +++++++ 1 file changed, 7 insertions(+) diff --git a/tensorrt_llm/usage/llm_args_golden_manifest.json b/tensorrt_llm/usage/llm_args_golden_manifest.json index 9f7ac4845d63..6b09d6d6e7a5 100644 --- a/tensorrt_llm/usage/llm_args_golden_manifest.json +++ b/tensorrt_llm/usage/llm_args_golden_manifest.json @@ -1825,6 +1825,13 @@ { "capture_policy": "bool", "kind": "value", + "path": "torch_compile_config.compile_only_piecewise_graphs" + }, + { + "allowed_values": [], + "annotation": "", + "converter": "", + "kind": "value", "path": "torch_compile_config.enable_fullgraph" }, { From f8992fb4cbf76e2c2429ee72f17d3fef3aa07bda Mon Sep 17 00:00:00 2001 From: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> Date: Mon, 7 Sep 2026 03:12:58 -0700 Subject: [PATCH 10/23] Add test for numerical stability of compile_only_piecewise_graphs option Signed-off-by: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> --- .../defs/accuracy/test_llm_api_pytorch.py | 50 +++++++++++++++++++ .../test_lists/qa/llm_function_core.txt | 1 + .../test_lists/test-db/l0_h100.yml | 1 + 3 files changed, 52 insertions(+) diff --git a/tests/integration/defs/accuracy/test_llm_api_pytorch.py b/tests/integration/defs/accuracy/test_llm_api_pytorch.py index d71b2c35e233..7b64ad91a6c6 100644 --- a/tests/integration/defs/accuracy/test_llm_api_pytorch.py +++ b/tests/integration/defs/accuracy/test_llm_api_pytorch.py @@ -2836,6 +2836,56 @@ def test_nvfp4(self, ep_size, attention_dp): task.evaluate(llm) +class TestQwen3_0_6B(LlmapiAccuracyTestHarness): + MODEL_NAME = "Qwen3/Qwen3-0.6B" + MODEL_PATH = f"{llm_models_root()}/Qwen3/Qwen3-0.6B" + + @skip_pre_hopper + @pytest.mark.timeout(1800) + def test_piecewise_cuda_graph_compile_restriction_stability(self) -> None: + """Compare compiled and eager decode after compiled piecewise prefill.""" + prompts = ["The capital of France is"] + sampling_params = SamplingParams(max_tokens=8, + end_id=-1, + temperature=0, + return_generation_logits=True) + + def run( + compile_only_piecewise_graphs: bool + ) -> tuple[list[int], torch.Tensor]: + torch_compile_config = TorchCompileConfig( + enable_fullgraph=True, + compile_only_piecewise_graphs=compile_only_piecewise_graphs, + ) + with LLM( + self.MODEL_PATH, + max_seq_len=128, + max_num_tokens=32, + max_batch_size=1, + disable_overlap_scheduler=True, + gather_generation_logits=True, + cuda_graph_config=CudaGraphConfig(enable_padding=True, + max_batch_size=1), + prefill_cuda_graph_backend=PrefillCudaGraphBackend. + PIECEWISE, + prefill_capture_num_tokens=[32], + torch_compile_config=torch_compile_config, + ) as llm: + output = llm.generate( + prompts, sampling_params=sampling_params)[0].outputs[0] + assert output.generation_logits is not None + return list(output.token_ids), output.generation_logits.clone() + + compiled_tokens, compiled_logits = run(False) + restricted_tokens, restricted_logits = run(True) + + assert restricted_tokens == compiled_tokens + torch.testing.assert_close(restricted_logits, + compiled_logits, + rtol=1e-2, + atol=1e-2) + + class TestQwen3_4B(LlmapiAccuracyTestHarness): MODEL_NAME = "Qwen3/Qwen3-4B" diff --git a/tests/integration/test_lists/qa/llm_function_core.txt b/tests/integration/test_lists/qa/llm_function_core.txt index 02b7f60e23bc..7921624db6f9 100644 --- a/tests/integration/test_lists/qa/llm_function_core.txt +++ b/tests/integration/test_lists/qa/llm_function_core.txt @@ -561,6 +561,7 @@ accuracy/test_llm_api_pytorch.py::TestQwen3_30B_A3B_Instruct_2507::test_skip_sof accuracy/test_llm_api_pytorch.py::TestQwen3_30B_A3B_Instruct_2507::test_skip_softmax_attention_4gpus[target_sparsity_0.5-fp8kv=True] accuracy/test_llm_api_pytorch.py::TestQwen3_30B_A3B_Instruct_2507::test_skip_softmax_attention_4gpus[target_sparsity_0.9-fp8kv=False] accuracy/test_llm_api_pytorch.py::TestQwen3_30B_A3B_Instruct_2507::test_skip_softmax_attention_4gpus[target_sparsity_0.9-fp8kv=True] +accuracy/test_llm_api_pytorch.py::TestQwen3_0_6B::test_piecewise_cuda_graph_compile_restriction_stability accuracy/test_llm_api_pytorch.py::TestQwen3_4B::test_eagle3 accuracy/test_llm_api_pytorch.py::TestQwen3_5_35B_A3B::test_bf16[tp1-CUTLASS] accuracy/test_llm_api_pytorch.py::TestQwen3_5_35B_A3B::test_bf16[tp1-TRTLLM] diff --git a/tests/integration/test_lists/test-db/l0_h100.yml b/tests/integration/test_lists/test-db/l0_h100.yml index 31659d64e83c..20d6030ef628 100644 --- a/tests/integration/test_lists/test-db/l0_h100.yml +++ b/tests/integration/test_lists/test-db/l0_h100.yml @@ -170,6 +170,7 @@ l0_h100: - accuracy/test_llm_api_pytorch.py::TestQwen3_8B::test_dflash - accuracy/test_llm_api_pytorch.py::TestQwen3_8B::test_dspark[VANILLA] - accuracy/test_llm_api_pytorch.py::TestGPTOSS::test_dflash + - accuracy/test_llm_api_pytorch.py::TestQwen3_0_6B::test_piecewise_cuda_graph_compile_restriction_stability - accuracy/test_llm_api_pytorch.py::TestQwen3_5_4B::test_bf16 - accuracy/test_llm_api_pytorch.py::TestQwen3_5_4B::test_fp8 - accuracy/test_llm_api_pytorch.py::TestQwen3_5_4B::test_fp8_piecewise_cuda_graph From c7fe471c59ce3512ddbb205288f8322a6111ac85 Mon Sep 17 00:00:00 2001 From: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> Date: Mon, 7 Sep 2026 04:51:07 -0700 Subject: [PATCH 11/23] Update docstring for compile_only_piecewise_graphs option Signed-off-by: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> --- tensorrt_llm/llmapi/llm_args.py | 9 +++++---- 1 file changed, 5 insertions(+), 4 deletions(-) diff --git a/tensorrt_llm/llmapi/llm_args.py b/tensorrt_llm/llmapi/llm_args.py index d4c2fbb87be4..0a0048a3f181 100644 --- a/tensorrt_llm/llmapi/llm_args.py +++ b/tensorrt_llm/llmapi/llm_args.py @@ -5765,10 +5765,11 @@ class TorchCompileConfig(StrictBaseModel): compile_only_piecewise_graphs: bool = Field( default=False, - description="Compile only context and mixed-batch execution graphs. " - "When enabled, generation-only and auxiliary kernel warmup forwards " - "use eager execution, while dedicated context/mixed compilation and " - "piecewise CUDA graph warmup and capture remain compiled.", + description="Compile only context and mixed-batch execution graphs, " + "while generation-only and auxiliary kernel warmup forwards " + "remain eager. Enabling this will speed up startup, but might " + "lead to degraded performance or inconsistent output in " + "specific configurations.", status="prototype") enable_inductor: bool = Field( From b76f082f39975191844c10109367c17ba3622da6 Mon Sep 17 00:00:00 2001 From: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> Date: Mon, 7 Sep 2026 05:54:38 -0700 Subject: [PATCH 12/23] Add checks for conflicting prefill backends and forbid with attention DP Signed-off-by: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> --- tensorrt_llm/llmapi/llm_args.py | 11 +++++++++++ 1 file changed, 11 insertions(+) diff --git a/tensorrt_llm/llmapi/llm_args.py b/tensorrt_llm/llmapi/llm_args.py index 0a0048a3f181..6678b0768a9f 100644 --- a/tensorrt_llm/llmapi/llm_args.py +++ b/tensorrt_llm/llmapi/llm_args.py @@ -6543,6 +6543,17 @@ def normalize_prefill_cuda_graph_config(self) -> 'TorchLlmArgs': if not buckets_are_explicit and legacy_buckets is not None: self.prefill_capture_num_tokens = list(legacy_buckets) + if (compile_config is not None + and compile_config.compile_only_piecewise_graphs): + if self.prefill_cuda_graph_backend != PrefillCudaGraphBackend.PIECEWISE: + raise ValueError( + "torch_compile_config.compile_only_piecewise_graphs requires " + "prefill_cuda_graph_backend='piecewise'") + if self.enable_attention_dp: + raise ValueError( + "torch_compile_config.compile_only_piecewise_graphs does not " + "support attention DP") + if self.prefill_cuda_graph_backend != PrefillCudaGraphBackend.DISABLED: if self.prefill_capture_num_tokens is None: self.prefill_capture_num_tokens = list( From 66c05c5d606a864eb9bf89b4c61b8972887f3db9 Mon Sep 17 00:00:00 2001 From: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> Date: Mon, 7 Sep 2026 05:55:24 -0700 Subject: [PATCH 13/23] Add tests for config validation Signed-off-by: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> --- tests/unittest/llmapi/test_llm_args.py | 17 +++++++++++++++++ 1 file changed, 17 insertions(+) diff --git a/tests/unittest/llmapi/test_llm_args.py b/tests/unittest/llmapi/test_llm_args.py index e95456b7272a..ee1f81cd5d49 100644 --- a/tests/unittest/llmapi/test_llm_args.py +++ b/tests/unittest/llmapi/test_llm_args.py @@ -2590,9 +2590,26 @@ def test_legacy_piecewise_config_maps_to_new_fields(self): args = TorchLlmArgs(model=llama_model_path, torch_compile_config=TorchCompileConfig( enable_piecewise_cuda_graph=True, + compile_only_piecewise_graphs=True, capture_num_tokens=[128, 256])) assert args.prefill_cuda_graph_backend == PrefillCudaGraphBackend.PIECEWISE assert args.prefill_capture_num_tokens == [256, 128] + assert args.torch_compile_config.compile_only_piecewise_graphs + + def test_compile_only_piecewise_graphs_requires_piecewise_backend(self): + with pytest.raises(ValueError, match="requires.*piecewise"): + TorchLlmArgs(model=llama_model_path, + torch_compile_config=TorchCompileConfig( + compile_only_piecewise_graphs=True)) + + def test_compile_only_piecewise_graphs_rejects_attention_dp(self): + with pytest.raises(ValueError, match="does not support attention DP"): + TorchLlmArgs( + model=llama_model_path, + enable_attention_dp=True, + prefill_cuda_graph_backend=PrefillCudaGraphBackend.PIECEWISE, + torch_compile_config=TorchCompileConfig( + compile_only_piecewise_graphs=True)) def test_explicit_new_buckets_with_legacy_piecewise_enable(self): args = TorchLlmArgs(model=llama_model_path, From 57f7eaea1ed455992174cb695bcb69322b5f8e02 Mon Sep 17 00:00:00 2001 From: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> Date: Mon, 7 Sep 2026 07:14:11 -0700 Subject: [PATCH 14/23] Move torch.compile bypass logic inside PyTorchModelEngine Signed-off-by: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> --- tensorrt_llm/_torch/compilation/utils.py | 24 ++------ .../_torch/pyexecutor/model_engine.py | 56 ++++++++++++------- tensorrt_llm/llmapi/llm_args.py | 10 ++-- .../test_phase_selective_forward.py | 41 ++++---------- 4 files changed, 56 insertions(+), 75 deletions(-) diff --git a/tensorrt_llm/_torch/compilation/utils.py b/tensorrt_llm/_torch/compilation/utils.py index 009e9d4fc07d..879ccb63a2fe 100644 --- a/tensorrt_llm/_torch/compilation/utils.py +++ b/tensorrt_llm/_torch/compilation/utils.py @@ -10,15 +10,15 @@ from torch.fx.experimental.symbolic_shapes import ShapeEnv from ..cuda_tile_utils import IS_CUDA_TILE_AVAILABLE -from ..utils import get_model_extra_attrs class _PhaseSelectiveForward: """ - This utility class is used to selectively bypass torch.compile - for operations that do not need it (attention warmup, auto-tuning), - as well as for generation-only forwards, where the ordinary CUDA - graph machinery is deemed sufficient. + This utility class is a proxy implementing an engine-controlled + torch.compile bypass, enabled by the compile_only_piecewise_graphs + option. This will then skip torch.compile for certain operations + (attention warmup, auto-tuning), as well as for generation forwards + and prefill/mixed forwards that are not graph-eligible. """ def __init__(self, eager_forward: Callable[..., object], @@ -28,19 +28,7 @@ def __init__(self, eager_forward: Callable[..., object], self._bypass_active = False def __call__(self, *args, **kwargs) -> object: - """ - The wrapper will enforce the eager forward function if: - 1) The torch compile bypass is active (e.g., for auto-tuning), or - 2) The batch is generation-only (rely on ordinary CUDA graphs) - """ - attrs = get_model_extra_attrs() - attn_metadata_ref = attrs.get("attention_metadata") if attrs else None - attn_metadata = (attn_metadata_ref() - if attn_metadata_ref is not None else None) - # This is the finalized batch classification after any - # context-to-generation promotion. - if (self._bypass_active or - (attn_metadata is not None and attn_metadata.num_contexts == 0)): + if self._bypass_active: return self._eager_forward(*args, **kwargs) return self._compiled_forward(*args, **kwargs) diff --git a/tensorrt_llm/_torch/pyexecutor/model_engine.py b/tensorrt_llm/_torch/pyexecutor/model_engine.py index 0198afdd84e7..58abf1c1dfae 100644 --- a/tensorrt_llm/_torch/pyexecutor/model_engine.py +++ b/tensorrt_llm/_torch/pyexecutor/model_engine.py @@ -653,7 +653,6 @@ def __init__( 'max_num_streams'].default self._torch_compile_enabled = torch_compile_enabled - self._compile_only_piecewise_graphs = compile_only_piecewise_graphs self._phase_selective_forward: Optional[_PhaseSelectiveForward] = None self._torch_compile_piecewise_cuda_graph = torch_compile_piecewise_cuda_graph self._torch_compile_prefill_only = False @@ -1299,13 +1298,11 @@ def _pad_batch_seed_mrope_delta_cache( request.py_mrope_delta_cache_slot = request.py_seq_slot @contextmanager - def _without_torch_compile(self): + def _maybe_bypass_torch_compile(self, bypass: bool = True): """ - When compile_only_piecewise_graphs is enabled in TorchCompileConfig, - this ctxt manager bypasses torch.compile invocation anywhere outside - of piecewise CUDA graph capture and warmup. + Bypass torch.compile for forwards in this context when requested. """ - if self._phase_selective_forward is None: + if self._phase_selective_forward is None or not bypass: yield return with self._phase_selective_forward.bypass(): @@ -1544,15 +1541,16 @@ def _warmup_scheduled(self, resource_manager: ResourceManager, self._prewarm_cute_dsl_indexer_q() log_mem_snapshot("warmup/after_cute_dsl_indexer_q") if not is_enc_dec: - with self._warmup_timer.phase( - "attention_jit", - metrics=self._metrics, - metric_name="attention_warmup_seconds"), self._without_torch_compile(): + with self._warmup_timer.phase("attention_jit", + metrics=self._metrics, + metric_name="attention_warmup_seconds" + ), self._maybe_bypass_torch_compile(): self._run_attention_warmup(resource_manager, can_run_general_warmup) if can_run_general_warmup: - # Specialize torch.compile graphs across the key input shapes before CUDA graph capture. + # Specialize torch.compile graphs for piecewise-graph-eligible + # input shapes before CUDA graph capture. with self._warmup_timer.phase("general", metrics=self._metrics, metric_name="general_warmup_seconds"): @@ -1560,8 +1558,8 @@ def _warmup_scheduled(self, resource_manager: ResourceManager, self._get_full_general_warmup_requests(resource_manager)) # Currently graph has not been captured, disable cuda graph for this warmup. with self.no_cuda_graph(): - # Do not bypass torch.compile in order to specialize the - # piecewise graphs before capture. + # Per-forward eligibility selects compiled execution for + # piecewise graph shapes and eager execution for graph misses. self._general_warmup(resource_manager, warmup_requests_configs) # Release C++ MoE workspace buffers so the autotuner can @@ -1576,10 +1574,10 @@ def _warmup_scheduled(self, resource_manager: ResourceManager, # Helix CP is decode-only and runs into issues with the # autotuner warmup's context requests. if not is_enc_dec and not self.mapping.has_cp_helix(): - with self._warmup_timer.phase( - "autotuner", - metrics=self._metrics, - metric_name="autotuner_warmup_seconds"), self._without_torch_compile(): + with self._warmup_timer.phase("autotuner", + metrics=self._metrics, + metric_name="autotuner_warmup_seconds" + ), self._maybe_bypass_torch_compile(): self._run_autotuner_warmup(resource_manager) log_mem_snapshot("warmup/after_autotuner") # Pre-JIT Mamba SSD multi-seq + HAS_INITSTATES=True Triton kernels @@ -1590,7 +1588,8 @@ def _warmup_scheduled(self, resource_manager: ResourceManager, with self._warmup_timer.phase( "mamba_hybrid", metrics=self._metrics, - metric_name="mamba_hybrid_warmup_seconds"), self._without_torch_compile(): + metric_name="mamba_hybrid_warmup_seconds" + ), self._maybe_bypass_torch_compile(): self._run_mamba_hybrid_warmup(resource_manager) log_mem_snapshot("warmup/after_mamba_hybrid") # Release the autotuner's exploration-mode intermediates. The @@ -1644,7 +1643,8 @@ def _warmup_scheduled(self, resource_manager: ResourceManager, with self._warmup_timer.phase( "memory_pool_prepop", metrics=self._metrics, - metric_name="memory_pool_prepopulation_seconds"), self._without_torch_compile(): + metric_name="memory_pool_prepopulation_seconds" + ), self._maybe_bypass_torch_compile(): warmup_requests_configs = self._get_max_shape_warmup_requests( resource_manager) self._general_warmup(resource_manager, warmup_requests_configs) @@ -6321,11 +6321,18 @@ def _forward_decoder( self._prepare_inputs_event.record() breakable_runner = self.breakable_cuda_graph_runner + # _prepare_inputs records the group-uniform prefill graph decision. + # An ordinary CUDA graph takes precedence; its capture must use the + # eager decoder when compilation is restricted to piecewise graphs. + use_compiled_forward = (not can_run_graph and + get_per_request_prefill_cuda_graph_flag()) with with_shared_pool(self.cuda_graph_runner.get_graph_pool()): def forward_step(): - with MoeLoadBalancerIterContext(moe_load_balancer): + with self._maybe_bypass_torch_compile( + bypass=not use_compiled_forward + ), MoeLoadBalancerIterContext(moe_load_balancer): return self._forward_step( inputs, gather_ids=gather_ids, @@ -6353,7 +6360,9 @@ def forward_step(): if needs_capture: def capture_forward_fn(inputs: Dict[str, Any]): - with MoeLoadBalancerIterContext(moe_load_balancer): + with self._maybe_bypass_torch_compile( + bypass=not use_compiled_forward + ), MoeLoadBalancerIterContext(moe_load_balancer): return self._forward_step( inputs, gather_ids=gather_ids, @@ -6399,6 +6408,11 @@ def capture_postprocess_fn(inputs: Dict[str, Any]): return outputs + def _runner_model_forward(self, **kwargs): + use_compiled_forward = get_per_request_prefill_cuda_graph_flag() + with self._maybe_bypass_torch_compile(bypass=not use_compiled_forward): + return self.model_forward(**kwargs) + def model_forward(self, **kwargs): assert self._model_caller is not None # Transitional: move this scope and reclaimer lifecycle into the decoder runner. diff --git a/tensorrt_llm/llmapi/llm_args.py b/tensorrt_llm/llmapi/llm_args.py index 6678b0768a9f..6e3d1ddb726f 100644 --- a/tensorrt_llm/llmapi/llm_args.py +++ b/tensorrt_llm/llmapi/llm_args.py @@ -5765,11 +5765,11 @@ class TorchCompileConfig(StrictBaseModel): compile_only_piecewise_graphs: bool = Field( default=False, - description="Compile only context and mixed-batch execution graphs, " - "while generation-only and auxiliary kernel warmup forwards " - "remain eager. Enabling this will speed up startup, but might " - "lead to degraded performance or inconsistent output in " - "specific configurations.", + description="Compile only forwards eligible for piecewise prefill " + "CUDA graphs, while generation-only forwards, prefill graph misses, " + "and auxiliary kernel warmup forwards remain eager. Enabling this " + "will speed up startup, but might lead to degraded performance or " + "inconsistent output in specific configurations.", status="prototype") enable_inductor: bool = Field( diff --git a/tests/unittest/_torch/compilation/test_phase_selective_forward.py b/tests/unittest/_torch/compilation/test_phase_selective_forward.py index d5f78b80ca4b..0b48ad7e6a10 100644 --- a/tests/unittest/_torch/compilation/test_phase_selective_forward.py +++ b/tests/unittest/_torch/compilation/test_phase_selective_forward.py @@ -1,48 +1,27 @@ # SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 -from types import SimpleNamespace -from unittest.mock import Mock, patch +from unittest.mock import Mock from tensorrt_llm._torch.compilation.utils import _PhaseSelectiveForward -def _attention_attrs(num_contexts: int) -> dict[str, object]: - metadata = SimpleNamespace(num_contexts=num_contexts) - return {"attention_metadata": lambda: metadata} - - -def test_phase_selective_forward_dispatches_by_batch_phase(): +def test_phase_selective_forward_uses_compiled_by_default() -> None: eager_forward = Mock(return_value="eager") compiled_forward = Mock(return_value="compiled") forward = _PhaseSelectiveForward(eager_forward, compiled_forward) - with patch( - "tensorrt_llm._torch.compilation.utils.get_model_extra_attrs", - return_value=_attention_attrs(num_contexts=1), - ): - assert forward("context") == "compiled" - - with patch( - "tensorrt_llm._torch.compilation.utils.get_model_extra_attrs", - return_value=_attention_attrs(num_contexts=0), - ): - assert forward("generation") == "eager" - - compiled_forward.assert_called_once_with("context") - eager_forward.assert_called_once_with("generation") + assert forward("input") == "compiled" + compiled_forward.assert_called_once_with("input") + eager_forward.assert_not_called() -def test_phase_selective_forward_bypass_is_restored(): +def test_phase_selective_forward_bypass_is_restored() -> None: eager_forward = Mock(return_value="eager") compiled_forward = Mock(return_value="compiled") forward = _PhaseSelectiveForward(eager_forward, compiled_forward) - with patch( - "tensorrt_llm._torch.compilation.utils.get_model_extra_attrs", - return_value=_attention_attrs(num_contexts=1), - ): - assert forward() == "compiled" - with forward.bypass(): - assert forward() == "eager" - assert forward() == "compiled" + assert forward() == "compiled" + with forward.bypass(): + assert forward() == "eager" + assert forward() == "compiled" From 6a53fd4187002a0389af933f10ffb8b928c1ac6e Mon Sep 17 00:00:00 2001 From: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> Date: Mon, 7 Sep 2026 08:23:44 -0700 Subject: [PATCH 15/23] Minor fix to model engine tests Signed-off-by: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> --- tests/unittest/_torch/executor/test_pytorch_model_engine.py | 1 + 1 file changed, 1 insertion(+) diff --git a/tests/unittest/_torch/executor/test_pytorch_model_engine.py b/tests/unittest/_torch/executor/test_pytorch_model_engine.py index d86a355f6d21..6fe87c5f3005 100644 --- a/tests/unittest/_torch/executor/test_pytorch_model_engine.py +++ b/tests/unittest/_torch/executor/test_pytorch_model_engine.py @@ -321,6 +321,7 @@ def _make_forward_only_engine( engine._fallback_to_engine = True engine._lora = SimpleNamespace(cuda_graph_manager=None) engine._force_lora_graph_for_capture = None + engine._phase_selective_forward = None semantic_attn_metadata = Mock() graph_attn_metadata = Mock() From 23f21d6fe27bed1e5a51352853efdd3edee2f28f Mon Sep 17 00:00:00 2001 From: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> Date: Wed, 9 Sep 2026 01:47:26 -0700 Subject: [PATCH 16/23] Encapsulate all torch.compile bypass logic inside _maybe_bypass_torch_compile Signed-off-by: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> --- .../_torch/pyexecutor/model_engine.py | 55 +++++++++++-------- 1 file changed, 33 insertions(+), 22 deletions(-) diff --git a/tensorrt_llm/_torch/pyexecutor/model_engine.py b/tensorrt_llm/_torch/pyexecutor/model_engine.py index 58abf1c1dfae..2f347508680b 100644 --- a/tensorrt_llm/_torch/pyexecutor/model_engine.py +++ b/tensorrt_llm/_torch/pyexecutor/model_engine.py @@ -1298,11 +1298,24 @@ def _pad_batch_seed_mrope_delta_cache( request.py_mrope_delta_cache_slot = request.py_seq_slot @contextmanager - def _maybe_bypass_torch_compile(self, bypass: bool = True): + def _maybe_bypass_torch_compile(self, + bypass: bool = False, + can_run_graph: bool = True): """ - Bypass torch.compile for forwards in this context when requested. + Bypass torch.compile explicitly or unless a piecewise graph can run. + + In particular, the bypass arg enforces an unconditional torch.compile + bypass; if this is False, can_run_graph then indicates whether there + was a CUDA graph key match, which means that the batch is generation-only + and can be skipped. If that's also False, then we check for piecewise + graph eligibility, and only in that case is torch.compile allowed to run. """ - if self._phase_selective_forward is None or not bypass: + if self._phase_selective_forward is None: + yield + return + bypass = (bypass or can_run_graph + or not get_per_request_prefill_cuda_graph_flag()) + if not bypass: yield return with self._phase_selective_forward.bypass(): @@ -1541,10 +1554,11 @@ def _warmup_scheduled(self, resource_manager: ResourceManager, self._prewarm_cute_dsl_indexer_q() log_mem_snapshot("warmup/after_cute_dsl_indexer_q") if not is_enc_dec: - with self._warmup_timer.phase("attention_jit", - metrics=self._metrics, - metric_name="attention_warmup_seconds" - ), self._maybe_bypass_torch_compile(): + with self._warmup_timer.phase( + "attention_jit", + metrics=self._metrics, + metric_name="attention_warmup_seconds"), self._maybe_bypass_torch_compile( + bypass=True): self._run_attention_warmup(resource_manager, can_run_general_warmup) @@ -1574,10 +1588,11 @@ def _warmup_scheduled(self, resource_manager: ResourceManager, # Helix CP is decode-only and runs into issues with the # autotuner warmup's context requests. if not is_enc_dec and not self.mapping.has_cp_helix(): - with self._warmup_timer.phase("autotuner", - metrics=self._metrics, - metric_name="autotuner_warmup_seconds" - ), self._maybe_bypass_torch_compile(): + with self._warmup_timer.phase( + "autotuner", + metrics=self._metrics, + metric_name="autotuner_warmup_seconds"), self._maybe_bypass_torch_compile( + bypass=True): self._run_autotuner_warmup(resource_manager) log_mem_snapshot("warmup/after_autotuner") # Pre-JIT Mamba SSD multi-seq + HAS_INITSTATES=True Triton kernels @@ -1588,8 +1603,8 @@ def _warmup_scheduled(self, resource_manager: ResourceManager, with self._warmup_timer.phase( "mamba_hybrid", metrics=self._metrics, - metric_name="mamba_hybrid_warmup_seconds" - ), self._maybe_bypass_torch_compile(): + metric_name="mamba_hybrid_warmup_seconds"), self._maybe_bypass_torch_compile( + bypass=True): self._run_mamba_hybrid_warmup(resource_manager) log_mem_snapshot("warmup/after_mamba_hybrid") # Release the autotuner's exploration-mode intermediates. The @@ -1643,8 +1658,8 @@ def _warmup_scheduled(self, resource_manager: ResourceManager, with self._warmup_timer.phase( "memory_pool_prepop", metrics=self._metrics, - metric_name="memory_pool_prepopulation_seconds" - ), self._maybe_bypass_torch_compile(): + metric_name="memory_pool_prepopulation_seconds"), self._maybe_bypass_torch_compile( + bypass=True): warmup_requests_configs = self._get_max_shape_warmup_requests( resource_manager) self._general_warmup(resource_manager, warmup_requests_configs) @@ -6324,14 +6339,11 @@ def _forward_decoder( # _prepare_inputs records the group-uniform prefill graph decision. # An ordinary CUDA graph takes precedence; its capture must use the # eager decoder when compilation is restricted to piecewise graphs. - use_compiled_forward = (not can_run_graph and - get_per_request_prefill_cuda_graph_flag()) - with with_shared_pool(self.cuda_graph_runner.get_graph_pool()): def forward_step(): with self._maybe_bypass_torch_compile( - bypass=not use_compiled_forward + can_run_graph=can_run_graph, ), MoeLoadBalancerIterContext(moe_load_balancer): return self._forward_step( inputs, @@ -6361,7 +6373,7 @@ def forward_step(): def capture_forward_fn(inputs: Dict[str, Any]): with self._maybe_bypass_torch_compile( - bypass=not use_compiled_forward + can_run_graph=can_run_graph, ), MoeLoadBalancerIterContext(moe_load_balancer): return self._forward_step( inputs, @@ -6409,8 +6421,7 @@ def capture_postprocess_fn(inputs: Dict[str, Any]): return outputs def _runner_model_forward(self, **kwargs): - use_compiled_forward = get_per_request_prefill_cuda_graph_flag() - with self._maybe_bypass_torch_compile(bypass=not use_compiled_forward): + with self._maybe_bypass_torch_compile(can_run_graph=False): return self.model_forward(**kwargs) def model_forward(self, **kwargs): From 5f6451f60438e357603e1d7f7dbbed0cce505672 Mon Sep 17 00:00:00 2001 From: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> Date: Wed, 9 Sep 2026 01:51:41 -0700 Subject: [PATCH 17/23] Add focused tests for graph eligibility, move stability test to unittests Signed-off-by: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> --- .../defs/accuracy/test_llm_api_pytorch.py | 50 -------- .../test_lists/qa/llm_function_core.txt | 1 - .../test_lists/test-db/l0_h100.yml | 1 - .../test_phase_selective_forward.py | 112 +++++++++++++++++- 4 files changed, 111 insertions(+), 53 deletions(-) diff --git a/tests/integration/defs/accuracy/test_llm_api_pytorch.py b/tests/integration/defs/accuracy/test_llm_api_pytorch.py index 7b64ad91a6c6..d71b2c35e233 100644 --- a/tests/integration/defs/accuracy/test_llm_api_pytorch.py +++ b/tests/integration/defs/accuracy/test_llm_api_pytorch.py @@ -2836,56 +2836,6 @@ def test_nvfp4(self, ep_size, attention_dp): task.evaluate(llm) -class TestQwen3_0_6B(LlmapiAccuracyTestHarness): - MODEL_NAME = "Qwen3/Qwen3-0.6B" - MODEL_PATH = f"{llm_models_root()}/Qwen3/Qwen3-0.6B" - - @skip_pre_hopper - @pytest.mark.timeout(1800) - def test_piecewise_cuda_graph_compile_restriction_stability(self) -> None: - """Compare compiled and eager decode after compiled piecewise prefill.""" - prompts = ["The capital of France is"] - sampling_params = SamplingParams(max_tokens=8, - end_id=-1, - temperature=0, - return_generation_logits=True) - - def run( - compile_only_piecewise_graphs: bool - ) -> tuple[list[int], torch.Tensor]: - torch_compile_config = TorchCompileConfig( - enable_fullgraph=True, - compile_only_piecewise_graphs=compile_only_piecewise_graphs, - ) - with LLM( - self.MODEL_PATH, - max_seq_len=128, - max_num_tokens=32, - max_batch_size=1, - disable_overlap_scheduler=True, - gather_generation_logits=True, - cuda_graph_config=CudaGraphConfig(enable_padding=True, - max_batch_size=1), - prefill_cuda_graph_backend=PrefillCudaGraphBackend. - PIECEWISE, - prefill_capture_num_tokens=[32], - torch_compile_config=torch_compile_config, - ) as llm: - output = llm.generate( - prompts, sampling_params=sampling_params)[0].outputs[0] - assert output.generation_logits is not None - return list(output.token_ids), output.generation_logits.clone() - - compiled_tokens, compiled_logits = run(False) - restricted_tokens, restricted_logits = run(True) - - assert restricted_tokens == compiled_tokens - torch.testing.assert_close(restricted_logits, - compiled_logits, - rtol=1e-2, - atol=1e-2) - - class TestQwen3_4B(LlmapiAccuracyTestHarness): MODEL_NAME = "Qwen3/Qwen3-4B" diff --git a/tests/integration/test_lists/qa/llm_function_core.txt b/tests/integration/test_lists/qa/llm_function_core.txt index 7921624db6f9..02b7f60e23bc 100644 --- a/tests/integration/test_lists/qa/llm_function_core.txt +++ b/tests/integration/test_lists/qa/llm_function_core.txt @@ -561,7 +561,6 @@ accuracy/test_llm_api_pytorch.py::TestQwen3_30B_A3B_Instruct_2507::test_skip_sof accuracy/test_llm_api_pytorch.py::TestQwen3_30B_A3B_Instruct_2507::test_skip_softmax_attention_4gpus[target_sparsity_0.5-fp8kv=True] accuracy/test_llm_api_pytorch.py::TestQwen3_30B_A3B_Instruct_2507::test_skip_softmax_attention_4gpus[target_sparsity_0.9-fp8kv=False] accuracy/test_llm_api_pytorch.py::TestQwen3_30B_A3B_Instruct_2507::test_skip_softmax_attention_4gpus[target_sparsity_0.9-fp8kv=True] -accuracy/test_llm_api_pytorch.py::TestQwen3_0_6B::test_piecewise_cuda_graph_compile_restriction_stability accuracy/test_llm_api_pytorch.py::TestQwen3_4B::test_eagle3 accuracy/test_llm_api_pytorch.py::TestQwen3_5_35B_A3B::test_bf16[tp1-CUTLASS] accuracy/test_llm_api_pytorch.py::TestQwen3_5_35B_A3B::test_bf16[tp1-TRTLLM] diff --git a/tests/integration/test_lists/test-db/l0_h100.yml b/tests/integration/test_lists/test-db/l0_h100.yml index 20d6030ef628..31659d64e83c 100644 --- a/tests/integration/test_lists/test-db/l0_h100.yml +++ b/tests/integration/test_lists/test-db/l0_h100.yml @@ -170,7 +170,6 @@ l0_h100: - accuracy/test_llm_api_pytorch.py::TestQwen3_8B::test_dflash - accuracy/test_llm_api_pytorch.py::TestQwen3_8B::test_dspark[VANILLA] - accuracy/test_llm_api_pytorch.py::TestGPTOSS::test_dflash - - accuracy/test_llm_api_pytorch.py::TestQwen3_0_6B::test_piecewise_cuda_graph_compile_restriction_stability - accuracy/test_llm_api_pytorch.py::TestQwen3_5_4B::test_bf16 - accuracy/test_llm_api_pytorch.py::TestQwen3_5_4B::test_fp8 - accuracy/test_llm_api_pytorch.py::TestQwen3_5_4B::test_fp8_piecewise_cuda_graph diff --git a/tests/unittest/_torch/compilation/test_phase_selective_forward.py b/tests/unittest/_torch/compilation/test_phase_selective_forward.py index 0b48ad7e6a10..069be9174b45 100644 --- a/tests/unittest/_torch/compilation/test_phase_selective_forward.py +++ b/tests/unittest/_torch/compilation/test_phase_selective_forward.py @@ -1,9 +1,17 @@ # SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 -from unittest.mock import Mock +from unittest.mock import Mock, patch +import pytest +import torch +from utils.llm_data import llm_models_root +from utils.util import skip_pre_hopper + +from tensorrt_llm import LLM, SamplingParams from tensorrt_llm._torch.compilation.utils import _PhaseSelectiveForward +from tensorrt_llm._torch.pyexecutor.model_engine import PyTorchModelEngine +from tensorrt_llm.llmapi import CudaGraphConfig, PrefillCudaGraphBackend, TorchCompileConfig def test_phase_selective_forward_uses_compiled_by_default() -> None: @@ -25,3 +33,105 @@ def test_phase_selective_forward_bypass_is_restored() -> None: with forward.bypass(): assert forward() == "eager" assert forward() == "compiled" + + +@pytest.mark.parametrize( + ("can_run_graph", "prefill_graph_eligible", "expected"), + [ + pytest.param(False, True, "compiled", id="pcg-eligible-prefill"), + pytest.param(False, False, "eager", id="pcg-ineligible-prefill"), + pytest.param(True, True, "eager", id="ordinary-generation-capture"), + pytest.param(False, False, "eager", id="generation-graph-miss"), + ], +) +def test_model_engine_selects_torch_compile_forward( + can_run_graph: bool, + prefill_graph_eligible: bool, + expected: str, +) -> None: + # A phase-selective proxy is installed by ModelEngine initialization when + # compile_only_piecewise_graphs is enabled. + engine = object.__new__(PyTorchModelEngine) + eager_forward = Mock(return_value="eager") + compiled_forward = Mock(return_value="compiled") + engine._phase_selective_forward = _PhaseSelectiveForward(eager_forward, compiled_forward) + + with patch( + "tensorrt_llm._torch.pyexecutor.model_engine.get_per_request_prefill_cuda_graph_flag", + return_value=prefill_graph_eligible, + ) as get_prefill_graph_flag: + with engine._maybe_bypass_torch_compile(can_run_graph=can_run_graph): + actual = engine._phase_selective_forward() + + assert actual == expected + selected_forward = compiled_forward if expected == "compiled" else eager_forward + other_forward = eager_forward if expected == "compiled" else compiled_forward + selected_forward.assert_called_once_with() + other_forward.assert_not_called() + if can_run_graph: + get_prefill_graph_flag.assert_not_called() + else: + get_prefill_graph_flag.assert_called_once_with() + + +def test_model_engine_explicit_bypass_skips_graph_eligibility() -> None: + engine = object.__new__(PyTorchModelEngine) + eager_forward = Mock(return_value="eager") + compiled_forward = Mock(return_value="compiled") + engine._phase_selective_forward = _PhaseSelectiveForward(eager_forward, compiled_forward) + + with patch( + "tensorrt_llm._torch.pyexecutor.model_engine.get_per_request_prefill_cuda_graph_flag", + return_value=True, + ) as get_prefill_graph_flag: + with engine._maybe_bypass_torch_compile(bypass=True, can_run_graph=False): + actual = engine._phase_selective_forward() + + assert actual == "eager" + eager_forward.assert_called_once_with() + compiled_forward.assert_not_called() + get_prefill_graph_flag.assert_not_called() + + +@skip_pre_hopper +@pytest.mark.timeout(1800) +def test_piecewise_cuda_graph_compile_restriction_stability() -> None: + """Compare compiled and eager decode after compiled piecewise prefill.""" + models_root = llm_models_root() + if models_root is None: + pytest.skip("LLM_MODELS_ROOT is not available") + model_path = models_root / "Qwen3/Qwen3-0.6B" + if not model_path.exists(): + pytest.skip(f"Model not found: {model_path}") + + prompts = ["The capital of France is"] + sampling_params = SamplingParams( + max_tokens=8, end_id=-1, temperature=0, return_generation_logits=True + ) + + def run(compile_only_piecewise_graphs: bool) -> tuple[list[int], torch.Tensor]: + torch_compile_config = TorchCompileConfig( + enable_fullgraph=True, + compile_only_piecewise_graphs=compile_only_piecewise_graphs, + ) + with LLM( + model_path, + max_seq_len=128, + max_num_tokens=32, + max_batch_size=1, + disable_overlap_scheduler=True, + gather_generation_logits=True, + cuda_graph_config=CudaGraphConfig(enable_padding=True, max_batch_size=1), + prefill_cuda_graph_backend=PrefillCudaGraphBackend.PIECEWISE, + prefill_capture_num_tokens=[32], + torch_compile_config=torch_compile_config, + ) as llm: + output = llm.generate(prompts, sampling_params=sampling_params)[0].outputs[0] + assert output.generation_logits is not None + return list(output.token_ids), output.generation_logits.clone() + + compiled_tokens, compiled_logits = run(False) + restricted_tokens, restricted_logits = run(True) + + assert restricted_tokens == compiled_tokens + torch.testing.assert_close(restricted_logits, compiled_logits, rtol=1e-2, atol=1e-2) From 57d658b87fe32ea8f80a00ad098e5616cec75c53 Mon Sep 17 00:00:00 2001 From: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> Date: Wed, 9 Sep 2026 02:11:47 -0700 Subject: [PATCH 18/23] Add documentation Signed-off-by: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> --- .../torch_compile_and_piecewise_cuda_graph.md | 19 +++++++++++++++++++ 1 file changed, 19 insertions(+) diff --git a/docs/source/features/torch_compile_and_piecewise_cuda_graph.md b/docs/source/features/torch_compile_and_piecewise_cuda_graph.md index 8b179030177f..ede783b765c0 100644 --- a/docs/source/features/torch_compile_and_piecewise_cuda_graph.md +++ b/docs/source/features/torch_compile_and_piecewise_cuda_graph.md @@ -13,6 +13,7 @@ Piecewise CUDA Graph is a technique that runs cudagraph-unsupported components ( - [Piecewise CUDA Graph & Generation Only CUDA Graph](#piecewise-cuda-graph--generation-only-cuda-graph) - [Piecewise CUDA Graph Padding](#piecewise-cuda-graph-padding) - [Performance Tuning](#performance-tuning) + - [Torch Compile Optimizations](#torch-compile-optimizations) - [Known Issue](#known-issue) - [Development Guide](#development-guide) - [Background Knowledge](#background-knowledge) @@ -113,6 +114,24 @@ Guidelines for `capture_num_tokens`: Even with Piecewise CUDA Graph enabled, you may still observe bubbles in the context (prefill) phase, primarily due to the attention operator’s substantial host-side overhead. +### Torch Compile Optimizations + +Enabling `prefill_cuda_graph_backend: piecewise` also enables `torch.compile`. By default, decoder forwards are routed through the compiled callable, including autotuning and auxiliary warmups as well as context, mixed, and generation-only batches. For some models, running these forwards through the compiled callable can improve performance through optimizations such as kernel fusion. In other cases, most of the benefit comes from applying piecewise CUDA graphs to context and mixed batches only. + +The `compile_only_piecewise_graphs` option restricts compilation to forwards that are eligible for piecewise CUDA graphs, including the corresponding specialization warmup and capture forwards: + +```yaml +prefill_cuda_graph_backend: piecewise +torch_compile_config: + compile_only_piecewise_graphs: true +``` + +With this option enabled, ineligible context and mixed forwards use the eager decoder. Ordinary CUDA graphs are also captured from the eager decoder and subsequently replayed as CUDA graphs. Auxiliary kernel and memory-pool warmups bypass `torch.compile`. Avoiding unnecessary tracing and compilation can significantly reduce startup time. + +The `compile_only_piecewise_graphs` option was validated with a Qwen3 8B FP8 model, where no performance impact was measured in the tested TP1 and TP2 configurations. However, performance and numerical behavior can be model- and configuration-dependent, so these should be evaluated on a case-by-case basis before deployment. +This option requires `prefill_cuda_graph_backend: piecewise`, does not support attention DP, and currently supports only models derived from `DecoderModelForCausalLM`. + + ## Known Issue Torch compile cannot work with multi-ModelEngine config, which currently means From 5b3bc431296514974a221c203e2d8723e46e1818 Mon Sep 17 00:00:00 2001 From: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> Date: Wed, 9 Sep 2026 02:23:08 -0700 Subject: [PATCH 19/23] Update comments after refactoring Signed-off-by: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> --- tensorrt_llm/_torch/pyexecutor/model_engine.py | 9 ++++----- 1 file changed, 4 insertions(+), 5 deletions(-) diff --git a/tensorrt_llm/_torch/pyexecutor/model_engine.py b/tensorrt_llm/_torch/pyexecutor/model_engine.py index 2f347508680b..8fab688a21b3 100644 --- a/tensorrt_llm/_torch/pyexecutor/model_engine.py +++ b/tensorrt_llm/_torch/pyexecutor/model_engine.py @@ -1572,8 +1572,6 @@ def _warmup_scheduled(self, resource_manager: ResourceManager, self._get_full_general_warmup_requests(resource_manager)) # Currently graph has not been captured, disable cuda graph for this warmup. with self.no_cuda_graph(): - # Per-forward eligibility selects compiled execution for - # piecewise graph shapes and eager execution for graph misses. self._general_warmup(resource_manager, warmup_requests_configs) # Release C++ MoE workspace buffers so the autotuner can @@ -6336,12 +6334,13 @@ def _forward_decoder( self._prepare_inputs_event.record() breakable_runner = self.breakable_cuda_graph_runner - # _prepare_inputs records the group-uniform prefill graph decision. - # An ordinary CUDA graph takes precedence; its capture must use the - # eager decoder when compilation is restricted to piecewise graphs. + with with_shared_pool(self.cuda_graph_runner.get_graph_pool()): def forward_step(): + # _prepare_inputs records the group-uniform prefill graph decision. + # An ordinary CUDA graph takes precedence; its capture must use the + # eager decoder when compilation is restricted to piecewise graphs. with self._maybe_bypass_torch_compile( can_run_graph=can_run_graph, ), MoeLoadBalancerIterContext(moe_load_balancer): From 249537cb74738423be44791b88d71b357ce377e6 Mon Sep 17 00:00:00 2001 From: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> Date: Thu, 17 Sep 2026 03:05:17 -0700 Subject: [PATCH 20/23] Fix stale LLM args manifest Signed-off-by: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> --- tensorrt_llm/usage/llm_args_golden_manifest.json | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/tensorrt_llm/usage/llm_args_golden_manifest.json b/tensorrt_llm/usage/llm_args_golden_manifest.json index 6b09d6d6e7a5..f1276afae720 100644 --- a/tensorrt_llm/usage/llm_args_golden_manifest.json +++ b/tensorrt_llm/usage/llm_args_golden_manifest.json @@ -1828,9 +1828,7 @@ "path": "torch_compile_config.compile_only_piecewise_graphs" }, { - "allowed_values": [], - "annotation": "", - "converter": "", + "capture_policy": "bool", "kind": "value", "path": "torch_compile_config.enable_fullgraph" }, From f72f471519d270d344231ed74df7007664abb754 Mon Sep 17 00:00:00 2001 From: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> Date: Mon, 28 Sep 2026 01:14:08 -0700 Subject: [PATCH 21/23] Integrate compile restriction with prefill wrapper Signed-off-by: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> --- .../torch_compile_and_piecewise_cuda_graph.md | 2 +- tensorrt_llm/_torch/compilation/utils.py | 35 +-------- .../_torch/pyexecutor/model_engine.py | 56 ++++++++------ tensorrt_llm/llmapi/llm_args.py | 4 - .../test_phase_selective_forward.py | 77 +++++++++++++------ .../executor/test_pytorch_model_engine.py | 3 +- .../test_pytorch_model_engine_warmup.py | 31 ++++++-- tests/unittest/llmapi/test_llm_args.py | 17 ++-- 8 files changed, 127 insertions(+), 98 deletions(-) diff --git a/docs/source/features/torch_compile_and_piecewise_cuda_graph.md b/docs/source/features/torch_compile_and_piecewise_cuda_graph.md index ede783b765c0..32ebadc25bea 100644 --- a/docs/source/features/torch_compile_and_piecewise_cuda_graph.md +++ b/docs/source/features/torch_compile_and_piecewise_cuda_graph.md @@ -129,7 +129,7 @@ torch_compile_config: With this option enabled, ineligible context and mixed forwards use the eager decoder. Ordinary CUDA graphs are also captured from the eager decoder and subsequently replayed as CUDA graphs. Auxiliary kernel and memory-pool warmups bypass `torch.compile`. Avoiding unnecessary tracing and compilation can significantly reduce startup time. The `compile_only_piecewise_graphs` option was validated with a Qwen3 8B FP8 model, where no performance impact was measured in the tested TP1 and TP2 configurations. However, performance and numerical behavior can be model- and configuration-dependent, so these should be evaluated on a case-by-case basis before deployment. -This option requires `prefill_cuda_graph_backend: piecewise`, does not support attention DP, and currently supports only models derived from `DecoderModelForCausalLM`. +This option requires `prefill_cuda_graph_backend: piecewise` and currently supports only models derived from `DecoderModelForCausalLM`. With attention DP, the compiled path follows the group-wide prefill graph eligibility decision. ## Known Issue diff --git a/tensorrt_llm/_torch/compilation/utils.py b/tensorrt_llm/_torch/compilation/utils.py index 879ccb63a2fe..2ec2cac1c3a1 100644 --- a/tensorrt_llm/_torch/compilation/utils.py +++ b/tensorrt_llm/_torch/compilation/utils.py @@ -2,7 +2,7 @@ # SPDX-License-Identifier: Apache-2.0 import contextlib -from collections.abc import Callable, Iterator +from collections.abc import Callable from typing import List, Optional, Union import torch @@ -12,39 +12,6 @@ from ..cuda_tile_utils import IS_CUDA_TILE_AVAILABLE -class _PhaseSelectiveForward: - """ - This utility class is a proxy implementing an engine-controlled - torch.compile bypass, enabled by the compile_only_piecewise_graphs - option. This will then skip torch.compile for certain operations - (attention warmup, auto-tuning), as well as for generation forwards - and prefill/mixed forwards that are not graph-eligible. - """ - - def __init__(self, eager_forward: Callable[..., object], - compiled_forward: Callable[..., object]) -> None: - self._eager_forward = eager_forward - self._compiled_forward = compiled_forward - self._bypass_active = False - - def __call__(self, *args, **kwargs) -> object: - if self._bypass_active: - return self._eager_forward(*args, **kwargs) - return self._compiled_forward(*args, **kwargs) - - @contextlib.contextmanager - def bypass(self) -> Iterator[None]: - """ - Disable torch.compile for all forwards under this ctxt manager. - """ - previous = self._bypass_active - self._bypass_active = True - try: - yield - finally: - self._bypass_active = previous - - def get_symint_val(i: Union[torch.SymInt | int]): if isinstance(i, int): return i diff --git a/tensorrt_llm/_torch/pyexecutor/model_engine.py b/tensorrt_llm/_torch/pyexecutor/model_engine.py index 8fab688a21b3..6850debce284 100644 --- a/tensorrt_llm/_torch/pyexecutor/model_engine.py +++ b/tensorrt_llm/_torch/pyexecutor/model_engine.py @@ -49,8 +49,7 @@ from ..autotuner import AutoTuner, autotune from ..compilation.backend import Backend from ..compilation.piecewise_optimizer import PiecewiseRunner -from ..compilation.utils import (_PhaseSelectiveForward, - capture_piecewise_cuda_graph) +from ..compilation.utils import capture_piecewise_cuda_graph from ..distributed import Distributed from ..distributed.communicator import init_pp_comm from ..memory_buffer_utils import clear_memory_buffers, with_shared_pool @@ -211,6 +210,7 @@ def __init__(self, eager_model: torch.nn.Module, """Keep eager and compiled entry points sharing the same model weights.""" super().__init__() self.eager_model = eager_model + self._bypass_active = False # The compiled callable references the same weights. Register only the # eager tree so state_dict(), children() and _apply() visit it once. object.__setattr__(self, "compiled_model", compiled_model) @@ -230,10 +230,24 @@ def named_modules( def forward(self, *args: Any, **kwargs: Any) -> Any: """Use the compiled path only for globally eligible prefill batches.""" model = (self.compiled_model - if get_per_request_prefill_cuda_graph_flag() else - self.eager_model) + if self.use_compiled_forward() else self.eager_model) return model(*args, **kwargs) + def use_compiled_forward(self) -> bool: + """Return whether the current forward should use torch.compile.""" + return (not self._bypass_active + and get_per_request_prefill_cuda_graph_flag()) + + @contextmanager + def bypass(self) -> Iterator[None]: + """Temporarily force eager execution regardless of graph eligibility.""" + previous = self._bypass_active + self._bypass_active = True + try: + yield + finally: + self._bypass_active = previous + def __getattr__(self, name: str) -> Any: """Delegate model-specific attributes to the original eager model.""" # Epilogues can access transformer attributes after forward returns. @@ -653,7 +667,8 @@ def __init__( 'max_num_streams'].default self._torch_compile_enabled = torch_compile_enabled - self._phase_selective_forward: Optional[_PhaseSelectiveForward] = None + self._compile_only_piecewise_graphs = compile_only_piecewise_graphs + self._prefill_compiled_model: Optional[_PrefillCompiledModel] = None self._torch_compile_piecewise_cuda_graph = torch_compile_piecewise_cuda_graph self._torch_compile_prefill_only = False @@ -719,9 +734,12 @@ def __init__( self._torch_compile_piecewise_cuda_graph and (compile_only_piecewise_graphs or not self.model.use_fx_for_pcg_fallback)) - self.model.model = ( - _PrefillCompiledModel(eager_model, compiled_model) - if self._torch_compile_prefill_only else compiled_model) + if self._torch_compile_prefill_only: + self._prefill_compiled_model = _PrefillCompiledModel( + eager_model, compiled_model) + self.model.model = self._prefill_compiled_model + else: + self.model.model = compiled_model elif callable(apply_llm_torch_compile): # TODO: Move this contract to MultimodalModelMixin once # multimodal models consistently expose their LLM compile @@ -1302,23 +1320,21 @@ def _maybe_bypass_torch_compile(self, bypass: bool = False, can_run_graph: bool = True): """ - Bypass torch.compile explicitly or unless a piecewise graph can run. + Bypass torch.compile explicitly or while an ordinary graph can run. - In particular, the bypass arg enforces an unconditional torch.compile - bypass; if this is False, can_run_graph then indicates whether there - was a CUDA graph key match, which means that the batch is generation-only - and can be skipped. If that's also False, then we check for piecewise - graph eligibility, and only in that case is torch.compile allowed to run. + The wrapper handles piecewise-graph eligibility directly. This context + additionally forces eager execution for auxiliary operations and + ordinary CUDA graph capture. """ - if self._phase_selective_forward is None: + if (not self._compile_only_piecewise_graphs + or self._prefill_compiled_model is None): yield return - bypass = (bypass or can_run_graph - or not get_per_request_prefill_cuda_graph_flag()) + bypass = bypass or can_run_graph if not bypass: yield return - with self._phase_selective_forward.bypass(): + with self._prefill_compiled_model.bypass(): yield @staticmethod @@ -6419,10 +6435,6 @@ def capture_postprocess_fn(inputs: Dict[str, Any]): return outputs - def _runner_model_forward(self, **kwargs): - with self._maybe_bypass_torch_compile(can_run_graph=False): - return self.model_forward(**kwargs) - def model_forward(self, **kwargs): assert self._model_caller is not None # Transitional: move this scope and reclaimer lifecycle into the decoder runner. diff --git a/tensorrt_llm/llmapi/llm_args.py b/tensorrt_llm/llmapi/llm_args.py index 6e3d1ddb726f..efb9aedb74b9 100644 --- a/tensorrt_llm/llmapi/llm_args.py +++ b/tensorrt_llm/llmapi/llm_args.py @@ -6549,10 +6549,6 @@ def normalize_prefill_cuda_graph_config(self) -> 'TorchLlmArgs': raise ValueError( "torch_compile_config.compile_only_piecewise_graphs requires " "prefill_cuda_graph_backend='piecewise'") - if self.enable_attention_dp: - raise ValueError( - "torch_compile_config.compile_only_piecewise_graphs does not " - "support attention DP") if self.prefill_cuda_graph_backend != PrefillCudaGraphBackend.DISABLED: if self.prefill_capture_num_tokens is None: diff --git a/tests/unittest/_torch/compilation/test_phase_selective_forward.py b/tests/unittest/_torch/compilation/test_phase_selective_forward.py index 069be9174b45..68777c857cd8 100644 --- a/tests/unittest/_torch/compilation/test_phase_selective_forward.py +++ b/tests/unittest/_torch/compilation/test_phase_selective_forward.py @@ -9,30 +9,45 @@ from utils.util import skip_pre_hopper from tensorrt_llm import LLM, SamplingParams -from tensorrt_llm._torch.compilation.utils import _PhaseSelectiveForward -from tensorrt_llm._torch.pyexecutor.model_engine import PyTorchModelEngine +from tensorrt_llm._torch.pyexecutor.model_engine import PyTorchModelEngine, _PrefillCompiledModel from tensorrt_llm.llmapi import CudaGraphConfig, PrefillCudaGraphBackend, TorchCompileConfig -def test_phase_selective_forward_uses_compiled_by_default() -> None: +def _create_prefill_compiled_model() -> tuple[_PrefillCompiledModel, Mock, Mock]: + eager_model = torch.nn.Module() + compiled_model = torch.nn.Module() eager_forward = Mock(return_value="eager") compiled_forward = Mock(return_value="compiled") - forward = _PhaseSelectiveForward(eager_forward, compiled_forward) + eager_model.forward = eager_forward + compiled_model.forward = compiled_forward + return (_PrefillCompiledModel(eager_model, compiled_model), eager_forward, compiled_forward) - assert forward("input") == "compiled" + +def test_prefill_compiled_model_uses_compiled_when_eligible() -> None: + model, eager_forward, compiled_forward = _create_prefill_compiled_model() + + with patch( + "tensorrt_llm._torch.pyexecutor.model_engine.get_per_request_prefill_cuda_graph_flag", + return_value=True, + ): + assert model("input") == "compiled" compiled_forward.assert_called_once_with("input") eager_forward.assert_not_called() -def test_phase_selective_forward_bypass_is_restored() -> None: - eager_forward = Mock(return_value="eager") - compiled_forward = Mock(return_value="compiled") - forward = _PhaseSelectiveForward(eager_forward, compiled_forward) +def test_prefill_compiled_model_bypass_is_restored() -> None: + model, eager_forward, compiled_forward = _create_prefill_compiled_model() - assert forward() == "compiled" - with forward.bypass(): - assert forward() == "eager" - assert forward() == "compiled" + with patch( + "tensorrt_llm._torch.pyexecutor.model_engine.get_per_request_prefill_cuda_graph_flag", + return_value=True, + ): + assert model() == "compiled" + with model.bypass(): + assert model() == "eager" + assert model() == "compiled" + assert eager_forward.call_count == 1 + assert compiled_forward.call_count == 2 @pytest.mark.parametrize( @@ -49,19 +64,17 @@ def test_model_engine_selects_torch_compile_forward( prefill_graph_eligible: bool, expected: str, ) -> None: - # A phase-selective proxy is installed by ModelEngine initialization when - # compile_only_piecewise_graphs is enabled. engine = object.__new__(PyTorchModelEngine) - eager_forward = Mock(return_value="eager") - compiled_forward = Mock(return_value="compiled") - engine._phase_selective_forward = _PhaseSelectiveForward(eager_forward, compiled_forward) + model, eager_forward, compiled_forward = _create_prefill_compiled_model() + engine._compile_only_piecewise_graphs = True + engine._prefill_compiled_model = model with patch( "tensorrt_llm._torch.pyexecutor.model_engine.get_per_request_prefill_cuda_graph_flag", return_value=prefill_graph_eligible, ) as get_prefill_graph_flag: with engine._maybe_bypass_torch_compile(can_run_graph=can_run_graph): - actual = engine._phase_selective_forward() + actual = model() assert actual == expected selected_forward = compiled_forward if expected == "compiled" else eager_forward @@ -76,16 +89,16 @@ def test_model_engine_selects_torch_compile_forward( def test_model_engine_explicit_bypass_skips_graph_eligibility() -> None: engine = object.__new__(PyTorchModelEngine) - eager_forward = Mock(return_value="eager") - compiled_forward = Mock(return_value="compiled") - engine._phase_selective_forward = _PhaseSelectiveForward(eager_forward, compiled_forward) + model, eager_forward, compiled_forward = _create_prefill_compiled_model() + engine._compile_only_piecewise_graphs = True + engine._prefill_compiled_model = model with patch( "tensorrt_llm._torch.pyexecutor.model_engine.get_per_request_prefill_cuda_graph_flag", return_value=True, ) as get_prefill_graph_flag: with engine._maybe_bypass_torch_compile(bypass=True, can_run_graph=False): - actual = engine._phase_selective_forward() + actual = model() assert actual == "eager" eager_forward.assert_called_once_with() @@ -93,6 +106,24 @@ def test_model_engine_explicit_bypass_skips_graph_eligibility() -> None: get_prefill_graph_flag.assert_not_called() +def test_model_engine_bypass_requires_compile_only_option() -> None: + engine = object.__new__(PyTorchModelEngine) + model, eager_forward, compiled_forward = _create_prefill_compiled_model() + engine._compile_only_piecewise_graphs = False + engine._prefill_compiled_model = model + + with patch( + "tensorrt_llm._torch.pyexecutor.model_engine.get_per_request_prefill_cuda_graph_flag", + return_value=True, + ): + with engine._maybe_bypass_torch_compile(bypass=True): + actual = model() + + assert actual == "compiled" + compiled_forward.assert_called_once_with() + eager_forward.assert_not_called() + + @skip_pre_hopper @pytest.mark.timeout(1800) def test_piecewise_cuda_graph_compile_restriction_stability() -> None: diff --git a/tests/unittest/_torch/executor/test_pytorch_model_engine.py b/tests/unittest/_torch/executor/test_pytorch_model_engine.py index 6fe87c5f3005..985e45225d2b 100644 --- a/tests/unittest/_torch/executor/test_pytorch_model_engine.py +++ b/tests/unittest/_torch/executor/test_pytorch_model_engine.py @@ -321,7 +321,8 @@ def _make_forward_only_engine( engine._fallback_to_engine = True engine._lora = SimpleNamespace(cuda_graph_manager=None) engine._force_lora_graph_for_capture = None - engine._phase_selective_forward = None + engine._compile_only_piecewise_graphs = False + engine._prefill_compiled_model = None semantic_attn_metadata = Mock() graph_attn_metadata = Mock() diff --git a/tests/unittest/_torch/executor/test_pytorch_model_engine_warmup.py b/tests/unittest/_torch/executor/test_pytorch_model_engine_warmup.py index 50c876092f12..c20301decd6b 100644 --- a/tests/unittest/_torch/executor/test_pytorch_model_engine_warmup.py +++ b/tests/unittest/_torch/executor/test_pytorch_model_engine_warmup.py @@ -55,11 +55,13 @@ @pytest.mark.cpu_only @pytest.mark.parametrize("model_kind", ["default", "m3", "m3_vl", "gdn", "mamba"]) @pytest.mark.parametrize("compile_enabled,piecewise", [(False, False), (True, False), (True, True)]) +@pytest.mark.parametrize("compile_only_piecewise_graphs", [False, True]) def test_pcg_fx_fallback_policy_is_model_specific( monkeypatch: pytest.MonkeyPatch, model_kind: str, compile_enabled: bool, piecewise: bool, + compile_only_piecewise_graphs: bool, ) -> None: """Exercise the real constructor's routing gate without loading weights or CUDA.""" from tensorrt_llm._torch.models.modeling_minimaxm3 import ( @@ -89,7 +91,11 @@ def test_pcg_fx_fallback_policy_is_model_specific( eager = torch.nn.Linear(4, 4) model.model = eager compile_config = SimpleNamespace( - enable_fullgraph=True, enable_inductor=False, enable_userbuffers=False, max_num_streams=1 + enable_fullgraph=True, + compile_only_piecewise_graphs=compile_only_piecewise_graphs, + enable_inductor=False, + enable_userbuffers=False, + max_num_streams=1, ) llm_args = SimpleNamespace( encode_only=False, @@ -164,7 +170,11 @@ def test_pcg_fx_fallback_policy_is_model_specific( checkpoint_loader=Mock(), ) - expected_prefill_only = compile_enabled and piecewise and model_kind in ("m3", "m3_vl") + expected_prefill_only = ( + compile_enabled + and piecewise + and (compile_only_piecewise_graphs or model_kind in ("m3", "m3_vl")) + ) assert engine._torch_compile_prefill_only is expected_prefill_only if not compile_enabled: compile_model.assert_not_called() @@ -267,11 +277,13 @@ def test_prefill_compile_preserves_partial_weight_reload( @pytest.mark.cpu_only @pytest.mark.parametrize("prefill_only", [False, True]) @pytest.mark.parametrize("eligible", [False, True]) +@pytest.mark.parametrize("bypass", [False, True]) @pytest.mark.parametrize("raises", [False, True]) def test_prefill_compile_scopes_whole_model_forward( monkeypatch: pytest.MonkeyPatch, prefill_only: bool, eligible: bool, + bypass: bool, raises: bool, ) -> None: """Keep compile state active through model epilogues and restore it on exit.""" @@ -286,7 +298,13 @@ def forward(**kwargs: object) -> str: observed.append(is_torch_compiling()) return "done" - model = SimpleNamespace(model_config=SimpleNamespace(extra_attrs={}), forward=forward) + eager_model = torch.nn.Identity() + eager_model.model_config = SimpleNamespace(extra_attrs={}) + eager_model.forward = forward + compiled_model = torch.nn.Identity() + compiled_model.forward = forward + prefill_compiled_model = _PrefillCompiledModel(eager_model, compiled_model) + model = prefill_compiled_model if prefill_only else eager_model engine = SimpleNamespace( _model_caller=ModelCaller(model, prefill_compile_only=prefill_only), _eager_workspace_reclaimer=None, @@ -297,14 +315,17 @@ def forward(**kwargs: object) -> str: model_call_module, "get_per_request_prefill_cuda_graph_flag", lambda: eligible ) monkeypatch.setattr(model_call_module, "is_trace_enabled", lambda name: False) - with torch_compiling(True): + bypass_scope = ( + prefill_compiled_model.bypass() if prefill_only and bypass else contextlib.nullcontext() + ) + with torch_compiling(True), bypass_scope: if raises: with pytest.raises(RuntimeError, match="epilogue failure"): PyTorchModelEngine.model_forward(engine, attn_metadata=Mock()) else: assert PyTorchModelEngine.model_forward(engine, attn_metadata=Mock()) == "done" assert is_torch_compiling() - expected = eligible if prefill_only else True + expected = eligible and not bypass if prefill_only else True assert observed == [expected] * (1 if raises else 2) diff --git a/tests/unittest/llmapi/test_llm_args.py b/tests/unittest/llmapi/test_llm_args.py index ee1f81cd5d49..1bf400a3aa95 100644 --- a/tests/unittest/llmapi/test_llm_args.py +++ b/tests/unittest/llmapi/test_llm_args.py @@ -2602,14 +2602,15 @@ def test_compile_only_piecewise_graphs_requires_piecewise_backend(self): torch_compile_config=TorchCompileConfig( compile_only_piecewise_graphs=True)) - def test_compile_only_piecewise_graphs_rejects_attention_dp(self): - with pytest.raises(ValueError, match="does not support attention DP"): - TorchLlmArgs( - model=llama_model_path, - enable_attention_dp=True, - prefill_cuda_graph_backend=PrefillCudaGraphBackend.PIECEWISE, - torch_compile_config=TorchCompileConfig( - compile_only_piecewise_graphs=True)) + def test_compile_only_piecewise_graphs_supports_attention_dp(self): + args = TorchLlmArgs( + model=llama_model_path, + enable_attention_dp=True, + prefill_cuda_graph_backend=PrefillCudaGraphBackend.PIECEWISE, + torch_compile_config=TorchCompileConfig( + compile_only_piecewise_graphs=True)) + assert args.enable_attention_dp + assert args.torch_compile_config.compile_only_piecewise_graphs def test_explicit_new_buckets_with_legacy_piecewise_enable(self): args = TorchLlmArgs(model=llama_model_path, From 2256e520449bb8b44f06a3d9cb530a5cd0447280 Mon Sep 17 00:00:00 2001 From: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> Date: Mon, 28 Sep 2026 01:20:35 -0700 Subject: [PATCH 22/23] Undo leftover changes in compile utils Signed-off-by: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> --- tensorrt_llm/_torch/compilation/utils.py | 6 +----- 1 file changed, 1 insertion(+), 5 deletions(-) diff --git a/tensorrt_llm/_torch/compilation/utils.py b/tensorrt_llm/_torch/compilation/utils.py index 2ec2cac1c3a1..2e98607b84b6 100644 --- a/tensorrt_llm/_torch/compilation/utils.py +++ b/tensorrt_llm/_torch/compilation/utils.py @@ -1,9 +1,5 @@ -# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 - import contextlib -from collections.abc import Callable -from typing import List, Optional, Union +from typing import Callable, List, Optional, Union import torch from torch.fx import Node From 8ee5a0f43a13fa247a12d8c6ccfcb9792d3093ff Mon Sep 17 00:00:00 2001 From: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> Date: Mon, 28 Sep 2026 06:52:13 -0700 Subject: [PATCH 23/23] Fix torch_compiling() assignment after rebase Signed-off-by: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> --- tensorrt_llm/_torch/pyexecutor/model_engine.py | 16 ++++++++-------- .../executor/test_pytorch_model_engine_warmup.py | 4 ++-- 2 files changed, 10 insertions(+), 10 deletions(-) diff --git a/tensorrt_llm/_torch/pyexecutor/model_engine.py b/tensorrt_llm/_torch/pyexecutor/model_engine.py index 6850debce284..47b96b38659c 100644 --- a/tensorrt_llm/_torch/pyexecutor/model_engine.py +++ b/tensorrt_llm/_torch/pyexecutor/model_engine.py @@ -1573,8 +1573,8 @@ def _warmup_scheduled(self, resource_manager: ResourceManager, with self._warmup_timer.phase( "attention_jit", metrics=self._metrics, - metric_name="attention_warmup_seconds"), self._maybe_bypass_torch_compile( - bypass=True): + metric_name="attention_warmup_seconds" + ), self._maybe_bypass_torch_compile(bypass=True): self._run_attention_warmup(resource_manager, can_run_general_warmup) @@ -1605,8 +1605,8 @@ def _warmup_scheduled(self, resource_manager: ResourceManager, with self._warmup_timer.phase( "autotuner", metrics=self._metrics, - metric_name="autotuner_warmup_seconds"), self._maybe_bypass_torch_compile( - bypass=True): + metric_name="autotuner_warmup_seconds" + ), self._maybe_bypass_torch_compile(bypass=True): self._run_autotuner_warmup(resource_manager) log_mem_snapshot("warmup/after_autotuner") # Pre-JIT Mamba SSD multi-seq + HAS_INITSTATES=True Triton kernels @@ -1617,8 +1617,8 @@ def _warmup_scheduled(self, resource_manager: ResourceManager, with self._warmup_timer.phase( "mamba_hybrid", metrics=self._metrics, - metric_name="mamba_hybrid_warmup_seconds"), self._maybe_bypass_torch_compile( - bypass=True): + metric_name="mamba_hybrid_warmup_seconds" + ), self._maybe_bypass_torch_compile(bypass=True): self._run_mamba_hybrid_warmup(resource_manager) log_mem_snapshot("warmup/after_mamba_hybrid") # Release the autotuner's exploration-mode intermediates. The @@ -1672,8 +1672,8 @@ def _warmup_scheduled(self, resource_manager: ResourceManager, with self._warmup_timer.phase( "memory_pool_prepop", metrics=self._metrics, - metric_name="memory_pool_prepopulation_seconds"), self._maybe_bypass_torch_compile( - bypass=True): + metric_name="memory_pool_prepopulation_seconds" + ), self._maybe_bypass_torch_compile(bypass=True): warmup_requests_configs = self._get_max_shape_warmup_requests( resource_manager) self._general_warmup(resource_manager, warmup_requests_configs) diff --git a/tests/unittest/_torch/executor/test_pytorch_model_engine_warmup.py b/tests/unittest/_torch/executor/test_pytorch_model_engine_warmup.py index c20301decd6b..c72a4df9667c 100644 --- a/tests/unittest/_torch/executor/test_pytorch_model_engine_warmup.py +++ b/tests/unittest/_torch/executor/test_pytorch_model_engine_warmup.py @@ -286,7 +286,7 @@ def test_prefill_compile_scopes_whole_model_forward( bypass: bool, raises: bool, ) -> None: - """Keep compile state active through model epilogues and restore it on exit.""" + """Scope compile dispatch by phase independently of the model wrapper.""" observed = [] def forward(**kwargs: object) -> str: @@ -325,7 +325,7 @@ def forward(**kwargs: object) -> str: else: assert PyTorchModelEngine.model_forward(engine, attn_metadata=Mock()) == "done" assert is_torch_compiling() - expected = eligible and not bypass if prefill_only else True + expected = eligible if prefill_only else True assert observed == [expected] * (1 if raises else 2)