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..32ebadc25bea 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` 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 diff --git a/tensorrt_llm/_torch/pyexecutor/model_engine.py b/tensorrt_llm/_torch/pyexecutor/model_engine.py index 2429877c9cb4..47b96b38659c 100644 --- a/tensorrt_llm/_torch/pyexecutor/model_engine.py +++ b/tensorrt_llm/_torch/pyexecutor/model_engine.py @@ -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) @@ -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. @@ -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 == @@ -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 @@ -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( @@ -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 @@ -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): """ @@ -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"): @@ -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 @@ -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 @@ -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) @@ -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, @@ -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, diff --git a/tensorrt_llm/llmapi/llm_args.py b/tensorrt_llm/llmapi/llm_args.py index 0e68df417237..efb9aedb74b9 100644 --- a/tensorrt_llm/llmapi/llm_args.py +++ b/tensorrt_llm/llmapi/llm_args.py @@ -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") + enable_inductor: bool = Field( default=False, description="Enable inductor backend in torch.compile.") @@ -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( diff --git a/tensorrt_llm/usage/llm_args_golden_manifest.json b/tensorrt_llm/usage/llm_args_golden_manifest.json index 9f7ac4845d63..f1276afae720 100644 --- a/tensorrt_llm/usage/llm_args_golden_manifest.json +++ b/tensorrt_llm/usage/llm_args_golden_manifest.json @@ -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", 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..68777c857cd8 --- /dev/null +++ b/tests/unittest/_torch/compilation/test_phase_selective_forward.py @@ -0,0 +1,168 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +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.pyexecutor.model_engine import PyTorchModelEngine, _PrefillCompiledModel +from tensorrt_llm.llmapi import CudaGraphConfig, PrefillCudaGraphBackend, TorchCompileConfig + + +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") + eager_model.forward = eager_forward + compiled_model.forward = compiled_forward + return (_PrefillCompiledModel(eager_model, compiled_model), eager_forward, compiled_forward) + + +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_prefill_compiled_model_bypass_is_restored() -> 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() == "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( + ("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: + engine = object.__new__(PyTorchModelEngine) + 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 = model() + + 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) + 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 = model() + + assert actual == "eager" + eager_forward.assert_called_once_with() + compiled_forward.assert_not_called() + 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: + """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) diff --git a/tests/unittest/_torch/executor/test_pytorch_model_engine.py b/tests/unittest/_torch/executor/test_pytorch_model_engine.py index d86a355f6d21..985e45225d2b 100644 --- a/tests/unittest/_torch/executor/test_pytorch_model_engine.py +++ b/tests/unittest/_torch/executor/test_pytorch_model_engine.py @@ -321,6 +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._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..c72a4df9667c 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,14 +277,16 @@ 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.""" + """Scope compile dispatch by phase independently of the model wrapper.""" observed = [] def forward(**kwargs: object) -> str: @@ -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,7 +315,10 @@ 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()) diff --git a/tests/unittest/llmapi/test_llm_args.py b/tests/unittest/llmapi/test_llm_args.py index e95456b7272a..1bf400a3aa95 100644 --- a/tests/unittest/llmapi/test_llm_args.py +++ b/tests/unittest/llmapi/test_llm_args.py @@ -2590,9 +2590,27 @@ 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_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,