diff --git a/tensorrt_llm/_torch/pyexecutor/model_engine.py b/tensorrt_llm/_torch/pyexecutor/model_engine.py index 86911732b02d..3ff126ba9e58 100644 --- a/tensorrt_llm/_torch/pyexecutor/model_engine.py +++ b/tensorrt_llm/_torch/pyexecutor/model_engine.py @@ -98,6 +98,46 @@ def warmup(self, resource_manager: ResourceManager) -> None: return +def _filter_piecewise_capture_num_tokens( + candidate_num_tokens: list[int], + max_num_tokens: int, + max_batch_size: int, + max_seq_len: int, + num_extra_decoding_steps: int = 0, +) -> Tuple[list[int], list[int]]: + """Cap piecewise CUDA graph capture candidates at the engine's reachable + `num_tokens` ceiling `max_batch_size * (max_seq_len - 1 - num_extra_decoding_steps)` + and ensure the ceiling itself is captured. + + Each in-flight request must leave room for at least one decode token, + so the ceiling is the largest forward-pass `num_tokens` the warmup + builder can construct. Including it in the capture set closes the + runtime padding gap between the next-largest candidate and the ceiling + (otherwise ISLs in that gap have no graph >= them and fall back to + eager). + + Returns `(kept, unrecordable)` where `kept` is sorted ascending, + deduped, and contains the ceiling whenever it is positive. + `unrecordable` is the sorted unique set of input entries above the + ceiling but within `max_num_tokens`. + """ + max_capturable_num_tokens = max( + 0, max_batch_size * (max_seq_len - 1 - num_extra_decoding_steps)) + piecewise_capacity_limit = min(max_num_tokens, max_capturable_num_tokens) + kept = sorted( + {i + for i in candidate_num_tokens if 0 < i <= piecewise_capacity_limit}) + if piecewise_capacity_limit > 0 and (not kept or kept[-1] + < piecewise_capacity_limit): + kept.append(piecewise_capacity_limit) + unrecordable = sorted({ + i + for i in candidate_num_tokens + if max_capturable_num_tokens < i <= max_num_tokens + }) + return kept, unrecordable + + def _filter_cuda_graph_batch_sizes(cuda_graph_batch_sizes: list[int], max_batch_size: int, max_num_tokens: int, max_total_draft_tokens: int, @@ -285,10 +325,23 @@ def __init__( torch_compile_piecewise_cuda_graph_num_tokens or cuda_graph_batch_sizes or []) - self._piecewise_cuda_graph_num_tokens = [ - i for i in piecewise_cuda_graph_num_tokens - if i <= self.max_num_tokens - ] + num_extra_decoding_steps = self._get_num_extra_decoding_steps() + self._piecewise_cuda_graph_num_tokens, unrecordable = ( + _filter_piecewise_capture_num_tokens( + piecewise_cuda_graph_num_tokens, + max_num_tokens=self.max_num_tokens, + max_batch_size=self.batch_size, + max_seq_len=self.max_seq_len, + num_extra_decoding_steps=num_extra_decoding_steps, + )) + if unrecordable: + logger.warning( + f"Skipping piecewise CUDA graph capture for num_tokens=" + f"{unrecordable}: exceeds reachable ceiling " + f"max_batch_size*(max_seq_len-1-num_extra_decoding_steps)=" + f"{max(0, self.batch_size * (self.max_seq_len - 1 - num_extra_decoding_steps))}. " + f"Capturing the ceiling itself; raise max_seq_len for larger graphs." + ) try: use_ub_for_nccl = ( diff --git a/tests/unittest/llmapi/test_llm_args.py b/tests/unittest/llmapi/test_llm_args.py index eee83183ced8..f05c90bcdb99 100644 --- a/tests/unittest/llmapi/test_llm_args.py +++ b/tests/unittest/llmapi/test_llm_args.py @@ -723,6 +723,228 @@ def test_generate_cuda_graph_batch_sizes_padding_edge_cases( assert max_batch_size in batch_sizes +class TestPiecewiseCudaGraphCaptureDefaults: + """Piecewise CUDA graph capture-set defaults and reachable-ceiling filter. + + Three invariants are exercised: + + 1. `TorchCompileConfig.capture_num_tokens` defaults to a fixed + powers-of-2 + 256-stride list when `enable_piecewise_cuda_graph` + is True (and stays `None` otherwise). The fixed list keeps the + capture set small to bound startup time and CUDA graph memory; + the model-engine filter (invariants 2 and 3) ensures the largest + reachable size is always captured even when it is not in this + default list. + 2. `_filter_piecewise_capture_num_tokens` caps the candidate list at + `max_batch_size * (max_seq_len - 1 - num_extra_decoding_steps)` -- + the largest forward-pass `num_tokens` the warmup builder can + construct, since every in-flight request must leave room for at + least one decode token. + 3. The reachable ceiling itself is always present in the returned + capture set (when positive), so runtime ISLs in the gap between + the next-largest candidate and the ceiling get a graph rather + than falling back to eager. + """ + + _EXPECTED_DEFAULT_CAPTURE_NUM_TOKENS = [2**i for i in range(8)] + list( + range(256, 3073, 256)) + + def test_torch_compile_config_capture_num_tokens_default_when_piecewise_enabled( + self): + """Default capture set is the powers-of-2 + 256-stride list. + + Keeps the capture set bounded (~20 entries) so server startup + time and CUDA graph memory stay predictable. The model engine + further filters and appends the reachable ceiling, so + out-of-range entries (e.g. > max_seq_len-1) are never recorded + and gap ISLs still get a graph. + """ + config = TorchCompileConfig(enable_piecewise_cuda_graph=True) + assert config.capture_num_tokens == self._EXPECTED_DEFAULT_CAPTURE_NUM_TOKENS + + def test_torch_compile_config_capture_num_tokens_stays_none_when_piecewise_disabled( + self): + """No default is populated when piecewise is off. + + The capture set is irrelevant in this case; populating it would + be misleading in serialized configs. + """ + config = TorchCompileConfig(enable_piecewise_cuda_graph=False) + assert config.capture_num_tokens is None + + def test_torch_compile_config_capture_num_tokens_user_override_preserved( + self): + """User-supplied `capture_num_tokens` is not overwritten by the default.""" + user_list = [4, 8, 16] + config = TorchCompileConfig(enable_piecewise_cuda_graph=True, + capture_num_tokens=user_list) + # `validate_capture_num_tokens` dedupes and reverse-sorts. + assert config.capture_num_tokens == sorted(set(user_list), reverse=True) + + def test_torch_llm_args_capture_num_tokens_default_when_piecewise_enabled( + self): + """Same default applies when reached through `TorchLlmArgs` construction. + + This is the path real users hit via `trtllm-serve` YAML. + """ + args = TorchLlmArgs( + model=llama_model_path, + max_batch_size=1, + max_seq_len=128, + max_beam_width=10, + enable_chunked_prefill=True, + cuda_graph_config=CudaGraphConfig(max_batch_size=128, + enable_padding=True), + torch_compile_config=TorchCompileConfig( + enable_piecewise_cuda_graph=True), + ) + assert args.torch_compile_config.capture_num_tokens == self._EXPECTED_DEFAULT_CAPTURE_NUM_TOKENS + + def test_piecewise_filter_drops_entries_above_reachable_ceiling(self): + """Drop candidates above `max_batch_size * (max_seq_len - 1)`. + + Without the cap, the warmup loop would silently skip these entries + and the outer padding logic would pad to a target with no captured + graph. They must be removed from `kept` and surfaced in + `unrecordable` so the warning fires. The ceiling itself is then + appended so ISLs in the gap still get a graph. + """ + from tensorrt_llm._torch.pyexecutor.model_engine import \ + _filter_piecewise_capture_num_tokens + + candidates = CudaGraphConfig._generate_cuda_graph_batch_sizes( + 128, enable_padding=True) + max_capturable = 1 * (128 - 1) + # Precondition: candidate list contains at least one entry above + # the reachable ceiling, otherwise the assertions below are vacuous. + assert any(i > max_capturable for i in candidates), ( + "Test precondition no longer holds: cuda_graph_batch_sizes " + f"for max_batch_size=128 no longer contains entries above " + f"{max_capturable}. Update this test if CudaGraphConfig " + "behavior changed.") + + kept, unrecordable = _filter_piecewise_capture_num_tokens( + candidates, + max_num_tokens=128, + max_batch_size=1, + max_seq_len=128, + ) + + assert kept[-1] == max_capturable # ceiling appended + assert 120 in kept # densely-packed entries below ceiling preserved + assert 128 not in kept + assert unrecordable == [128] + + def test_piecewise_filter_keeps_all_entries_when_within_ceiling(self): + """Keep all candidates when the largest fits within the ceiling. + + Symmetric case: when `max_batch_size * (max_seq_len - 1)` is at + least as large as the biggest candidate, nothing is dropped and + `unrecordable` is empty. The ceiling (128 here) coincides with the + largest candidate so no extra entry is appended. + """ + from tensorrt_llm._torch.pyexecutor.model_engine import \ + _filter_piecewise_capture_num_tokens + + candidates = CudaGraphConfig._generate_cuda_graph_batch_sizes( + 128, enable_padding=True) + kept, unrecordable = _filter_piecewise_capture_num_tokens( + candidates, + max_num_tokens=129, + max_batch_size=1, + max_seq_len=129, + ) + assert kept[-1] == 128 + assert 128 in kept + # Ceiling (128) was already in candidates, must not be duplicated. + assert kept.count(128) == 1 + assert unrecordable == [] + + def test_piecewise_filter_subtracts_extra_decoding_steps(self): + """Subtract `num_extra_decoding_steps` from the ceiling. + + Drafting loops consume extra decode steps; the filter must mirror + the `max_seq_len - 1 - num_extra_decoding_steps` constraint + applied when warmup requests are built. The ceiling is appended + whenever it is strictly greater than the largest surviving + candidate. + """ + from tensorrt_llm._torch.pyexecutor.model_engine import \ + _filter_piecewise_capture_num_tokens + + candidates = [1, 2, 4, 8, 16, 32, 64, 100, 120] + # max_seq_len=128, batch=1, 5 extra decoding steps -> ceiling 122. + kept, unrecordable = _filter_piecewise_capture_num_tokens( + candidates, + max_num_tokens=128, + max_batch_size=1, + max_seq_len=128, + num_extra_decoding_steps=5, + ) + assert kept[-1] == 122 + assert 120 in kept + assert unrecordable == [] + # Same setup with 9 extra decoding steps -> ceiling 118; 120 drops. + kept, unrecordable = _filter_piecewise_capture_num_tokens( + candidates, + max_num_tokens=128, + max_batch_size=1, + max_seq_len=128, + num_extra_decoding_steps=9, + ) + assert kept[-1] == 118 + assert 100 in kept + assert 120 not in kept + assert unrecordable == [120] + + def test_piecewise_filter_does_not_double_append_ceiling(self): + """Ceiling already present in candidates -> not duplicated.""" + from tensorrt_llm._torch.pyexecutor.model_engine import \ + _filter_piecewise_capture_num_tokens + + kept, _ = _filter_piecewise_capture_num_tokens( + [1, 64, 128], + max_num_tokens=129, + max_batch_size=1, + max_seq_len=129, + ) + assert kept == [1, 64, 128] + + def test_piecewise_filter_returns_empty_when_ceiling_is_zero(self): + """`max_seq_len=1` -> ceiling 0 -> nothing captured. + + With ceiling 0 every positive candidate is unrecordable and the + ceiling itself is not appended, so the warning "exceeds reachable + ceiling 0; raise max_seq_len" fires for the full candidate list. + """ + from tensorrt_llm._torch.pyexecutor.model_engine import \ + _filter_piecewise_capture_num_tokens + + kept, unrecordable = _filter_piecewise_capture_num_tokens( + [1, 2, 4], + max_num_tokens=8, + max_batch_size=1, + max_seq_len=1, + ) + assert kept == [] + assert unrecordable == [1, 2, 4] + + def test_piecewise_filter_appends_ceiling_when_only_smaller_candidates( + self): + """No candidate near the ceiling -> ceiling still appended.""" + from tensorrt_llm._torch.pyexecutor.model_engine import \ + _filter_piecewise_capture_num_tokens + + kept, _ = _filter_piecewise_capture_num_tokens( + [1, 2, 4, 8], + max_num_tokens=1024, + max_batch_size=8, + max_seq_len=128, + ) + # Ceiling: 8 * (128 - 1) = 1016. + assert kept == [1, 2, 4, 8, 1016] + + class TestTrtLlmArgs: def test_dynamic_setattr(self): @@ -1356,7 +1578,6 @@ def _get_all_llm_args_classes(): def _get_all_pydantic_models_from_llm_args(): """Get all Pydantic models referenced by BaseLlmArgs and its subclasses.""" - visited = set() models = [] @@ -1409,8 +1630,7 @@ def _get_qualified_name(cls: type) -> str: class TestPydanticBestPractices: - """ - Ensure that the user-facing LlmArgs and its subfields follow Pydantic best practices. + """Ensure that the user-facing LlmArgs and its subfields follow Pydantic best practices. """ # Fields exempt from Pydantic compatibility checks due to typing limitations or other edge cases. @@ -1444,8 +1664,7 @@ class TestPydanticBestPractices: def _is_allowed_type(self, annotation, model_cls: type, field_name: str) -> tuple[bool, str]: - """ - Check if a type annotation is allowed for user-facing config fields. + """Check if a type annotation is allowed for user-facing config fields. Allowed: - Pydantic models (must inherit from StrictBaseModel) @@ -1457,7 +1676,6 @@ def _is_allowed_type(self, annotation, model_cls: type, Returns (is_allowed, reason) tuple. """ - # Check if this field is exempt (check class and all parent classes) for cls in model_cls.__mro__: if field_name in self._COMPATIBILITY_EXEMPT_FIELDS.get(cls, []): @@ -1543,7 +1761,7 @@ def test_all_fields_have_descriptions(self): if violations: pytest.fail( - f"The following fields are missing descriptions:\n" + + "The following fields are missing descriptions:\n" + "\n".join(violations) + "\n\nPlease add a description to each by using Field(description=\"...\")." ) @@ -1569,7 +1787,7 @@ def test_all_fields_have_allowed_types(self): if violations: pytest.fail( - f"The following user-facing fields have types that are not allowed:\n" + "The following user-facing fields have types that are not allowed:\n" + "\n".join(violations) + "\n\nPlease use Pydantic-compatible types (primitives, Pydantic models that inherit from StrictBaseModel, " "or other compatible types). If this is intentional, add the field " @@ -1612,7 +1830,7 @@ def test_no_manual_serialization_or_validation_methods(self): if violations: pytest.fail( - f"The following models define forbidden methods:\n" + + "The following models define forbidden methods:\n" + "\n".join(violations) + "\n\nPydantic models should follow the recommendations above instead." ) @@ -1635,7 +1853,7 @@ def test_no_mutable_default_values(self): if violations: pytest.fail( - f"The following fields use mutable default values:\n" + + "The following fields use mutable default values:\n" + "\n".join(violations) + "\n\nMutable defaults are shared across instances and cause bugs. " "Use Field(default_factory=...) instead.") @@ -1662,7 +1880,7 @@ def test_no_custom_init_methods(self): if violations: pytest.fail( - f"The following Pydantic models define custom __init__ methods:\n" + "The following Pydantic models define custom __init__ methods:\n" + "\n".join(violations) + "\n\nThese should be replaced with alternatives like validators, model_post_init, or classmethods. See this test's docstring for more details." )