Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
23 commits
Select commit Hold shift + click to select a range
2364066
[WiP] Allow disabling torch.compile for generation-only batches
AlessioNetti Aug 27, 2026
ee82d07
Refactor compile_generation flag to cover warmup and autotuning passes
AlessioNetti Aug 28, 2026
82fe10a
Rename new flag to compile_only_piecewise_graphs
AlessioNetti Aug 31, 2026
3cfea35
Minor refactoring
AlessioNetti Aug 31, 2026
7ccba0d
Encapsulate torch.compile bypass logic in helper class
AlessioNetti Aug 31, 2026
c740f37
Drop trivial config tests
AlessioNetti Aug 31, 2026
70bd3bd
Add documentation
AlessioNetti Aug 31, 2026
8ec255a
Add comments to justify missing torch.compile bypass
AlessioNetti Sep 1, 2026
6f2f62f
Update LLM args manifest
AlessioNetti Sep 1, 2026
f8992fb
Add test for numerical stability of compile_only_piecewise_graphs option
AlessioNetti Sep 7, 2026
c7fe471
Update docstring for compile_only_piecewise_graphs option
AlessioNetti Sep 7, 2026
b76f082
Add checks for conflicting prefill backends and forbid with attention DP
AlessioNetti Sep 7, 2026
66c05c5
Add tests for config validation
AlessioNetti Sep 7, 2026
57f7eae
Move torch.compile bypass logic inside PyTorchModelEngine
AlessioNetti Sep 7, 2026
6a53fd4
Minor fix to model engine tests
AlessioNetti Sep 7, 2026
23f21d6
Encapsulate all torch.compile bypass logic inside _maybe_bypass_torch…
AlessioNetti Sep 9, 2026
5f6451f
Add focused tests for graph eligibility, move stability test to unitt…
AlessioNetti Sep 9, 2026
57d658b
Add documentation
AlessioNetti Sep 9, 2026
5b3bc43
Update comments after refactoring
AlessioNetti Sep 9, 2026
249537c
Fix stale LLM args manifest
AlessioNetti Sep 17, 2026
f72f471
Integrate compile restriction with prefill wrapper
AlessioNetti Sep 28, 2026
2256e52
Undo leftover changes in compile utils
AlessioNetti Sep 28, 2026
8ee5a0f
Fix torch_compiling() assignment after rebase
AlessioNetti Sep 28, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
19 changes: 19 additions & 0 deletions docs/source/features/torch_compile_and_piecewise_cuda_graph.md
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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` and currently supports only models derived from `DecoderModelForCausalLM`. With attention DP, the compiled path follows the group-wide prefill graph eligibility decision.


## Known Issue

Torch compile cannot work with multi-ModelEngine config, which currently means
Expand Down
96 changes: 82 additions & 14 deletions tensorrt_llm/_torch/pyexecutor/model_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -210,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)
Expand All @@ -229,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.
Expand Down Expand Up @@ -638,6 +653,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_piecewise_graphs = (
self.torch_compile_config.compile_only_piecewise_graphs
if self.torch_compile_config is not None else TorchCompileConfig.
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 ==
Expand All @@ -648,6 +667,8 @@ def __init__(
'max_num_streams'].default

self._torch_compile_enabled = torch_compile_enabled
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

Expand Down Expand Up @@ -695,6 +716,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(
Expand All @@ -703,10 +732,14 @@ 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)
self.model.model = (
_PrefillCompiledModel(eager_model, compiled_model)
if self._torch_compile_prefill_only else compiled_model)
and (compile_only_piecewise_graphs
or not self.model.use_fx_for_pcg_fallback))
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
Expand Down Expand Up @@ -1282,6 +1315,28 @@ 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 _maybe_bypass_torch_compile(self,
bypass: bool = False,
can_run_graph: bool = True):
"""
Bypass torch.compile explicitly or while an ordinary graph can run.

The wrapper handles piecewise-graph eligibility directly. This context
additionally forces eager execution for auxiliary operations and
ordinary CUDA graph capture.
"""
if (not self._compile_only_piecewise_graphs
or self._prefill_compiled_model is None):
yield
return
bypass = bypass or can_run_graph
if not bypass:
yield
return
with self._prefill_compiled_model.bypass():
yield

@staticmethod
def warmup_with_kv_cache_cleanup(method):
"""
Expand Down Expand Up @@ -1518,12 +1573,14 @@ 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._maybe_bypass_torch_compile(bypass=True):
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"):
Expand All @@ -1548,7 +1605,8 @@ 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._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
Expand All @@ -1559,7 +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"):
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
Expand Down Expand Up @@ -1608,11 +1667,13 @@ 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,
metric_name="memory_pool_prepopulation_seconds"):
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)
Expand Down Expand Up @@ -6293,7 +6354,12 @@ def _forward_decoder(
with with_shared_pool(self.cuda_graph_runner.get_graph_pool()):

def forward_step():
with MoeLoadBalancerIterContext(moe_load_balancer):
# _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):
return self._forward_step(
inputs,
gather_ids=gather_ids,
Expand Down Expand Up @@ -6321,7 +6387,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(
can_run_graph=can_run_graph,
), MoeLoadBalancerIterContext(moe_load_balancer):
return self._forward_step(
inputs,
gather_ids=gather_ids,
Expand Down
16 changes: 16 additions & 0 deletions tensorrt_llm/llmapi/llm_args.py
Original file line number Diff line number Diff line change
Expand Up @@ -5763,6 +5763,15 @@ class TorchCompileConfig(StrictBaseModel):
default=True,
description="Enable full graph compilation in torch.compile.")

compile_only_piecewise_graphs: bool = Field(
default=False,
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")
Comment on lines +5766 to +5773

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

compile_only_piecewise_graphs=True is currently accepted while prefill_cuda_graph_backend remains disabled. In that configuration the backend is not piecewise, but every context/mixed forward still enters torch.compile, so the option silently violates its name and PR contract.

We should either require PrefillCudaGraphBackend.PIECEWISE in normalize_prefill_cuda_graph_config() and dispatch on actual piecewise eligibility, or rename the option to describe phase-selective compilation.

We should also add validation tests for disabled and conflicting backends.

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

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

Thanks - I had missed this. I added a check during config validation with commit 39249ef to prevent this, and also forbid using the new option under attention DP (see other comment). I also added simple validation tests with commit 08f5a04.


enable_inductor: bool = Field(
default=False, description="Enable inductor backend in torch.compile.")

Expand Down Expand Up @@ -6534,6 +6543,13 @@ 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.prefill_cuda_graph_backend != PrefillCudaGraphBackend.DISABLED:
if self.prefill_capture_num_tokens is None:
self.prefill_capture_num_tokens = list(
Expand Down
5 changes: 5 additions & 0 deletions tensorrt_llm/usage/llm_args_golden_manifest.json
Original file line number Diff line number Diff line change
Expand Up @@ -1822,6 +1822,11 @@
"kind": "value",
"path": "torch_compile_config.capture_num_tokens"
},
{
"capture_policy": "bool",
"kind": "value",
"path": "torch_compile_config.compile_only_piecewise_graphs"
},
{
"capture_policy": "bool",
"kind": "value",
Expand Down
Loading
Loading