diff --git a/tensorrt_llm/_torch/peft/lora/layer.py b/tensorrt_llm/_torch/peft/lora/layer.py index f26ea410b535..4d7e87de1935 100644 --- a/tensorrt_llm/_torch/peft/lora/layer.py +++ b/tensorrt_llm/_torch/peft/lora/layer.py @@ -14,14 +14,19 @@ # limitations under the License. from collections.abc import Callable +from copy import copy from dataclasses import dataclass from enum import IntEnum from typing import Dict, List, Optional import torch +from ...autotuner import (AutoTuner, DynamicTensorSpec, OptimizationProfile, + TunableRunner, TuningConfig) from ...modules.multi_stream_utils import (do_multi_stream, maybe_execute_in_parallel) +from ...utils import (get_last_power_of_2_num_tokens_buckets, + last_positive_power_of_2) from .cuda_graph_lora_params import CudaGraphLoraParams _FP8_LORA_TMA_ALIGNMENT = 16 @@ -64,6 +69,10 @@ def _validate_fp8_lora_cuda_graph_alignment(slot_ranks_host: torch.Tensor, return min(hidden_size, min_active_rank) +_LORA_DEFAULT_SPLIT_K = 16 +_LORA_SPLIT_K_CANDIDATES = (1, 2, 4, 8, 16) + + @dataclass class GroupedGemmParamsOutput: in_sizes: Optional[torch.Tensor] = None @@ -232,6 +241,7 @@ def __init__(self, lora_module_types: List[LoraModuleType], assert len(lora_module_types) == len(output_hidden_sizes) self._par_events: List[torch.cuda.Event] | None = None + self._split_k_runner: Optional["_LoraGroupedGemmRunner"] = None @staticmethod def forward_with_base( @@ -549,19 +559,21 @@ def _prepare_max_sizes_cpu(self, return host_max_in_sizes, host_max_out_sizes - def _forward_cuda_graph_mode( + def _forward_cuda_graph_mode_impl( self, x: torch.Tensor, lora_params: Dict, layer_idx: int, + split_k: int, ) -> Optional[torch.Tensor]: """ - Forward pass using CUDA Graph compatible LoRA parameters. + Run the complete CUDA-graph LoRA path with a fixed split-K. Args: x: Input tensor lora_params: CUDA Graph compatible LoRA parameters layer_idx: Current layer index + split_k: Fixed split-K value chosen by the autotuner Returns: LoRA output tensor or None @@ -636,7 +648,7 @@ def _forward_cuda_graph_mode( grouped_gemm_params.ldd, grouped_gemm_params.ldb_prime, grouped_gemm_params.ldd_prime, host_max_in_sizes, host_max_out_sizes, grouped_gemm_params.splitk_offsets, - grouped_gemm_params.reordered_input.dtype, min_kn) + grouped_gemm_params.reordered_input.dtype, min_kn, split_k) # PyTorch does not implement index_copy_ for FP8 tensors. if output_buffer.dtype == torch.float8_e4m3fn: @@ -650,6 +662,59 @@ def _forward_cuda_graph_mode( output_buffer) return restored_output + def _forward_cuda_graph_mode( + self, + x: torch.Tensor, + lora_params: Dict, + layer_idx: int, + ) -> Optional[torch.Tensor]: + """ + Forward pass using CUDA Graph compatible LoRA parameters. + + Args: + x: Input tensor + lora_params: CUDA Graph compatible LoRA parameters + layer_idx: Current layer index + + Returns: + LoRA output tensor or None + """ + + cuda_graph_params: CudaGraphLoraParams = lora_params.get( + 'cuda_graph_params') + # Get layer-specific parameters + layer_key = CudaGraphLoraParams.LoraLayerKey( + layer_idx=layer_idx, module_ids=tuple(self.lora_module_types)) + + if not cuda_graph_params or not cuda_graph_params.layer_info or layer_key not in cuda_graph_params.layer_info: + return None + + # Skip layers that don't have LoRA modules + layer_params = cuda_graph_params.get_layer_params(layer_key) + if layer_params is None: + return None # Pass-through for layers without LoRA modules + if self._split_k_runner is None: + self._split_k_runner = _LoraGroupedGemmRunner( + layer=self, + layer_idx=layer_idx, + input_hidden_size=x.shape[-1], + max_rank=cuda_graph_params.max_rank, + max_lora_size=cuda_graph_params.max_lora_size, + problem_count=cuda_graph_params.get_problem_count(layer_key), + dtype=x.dtype, + ) + + runner = self._split_k_runner + runner.lora_params = runner.copy_lora_params(lora_params) + runner_inputs = [x] + _, split_k = AutoTuner.get().choose_one( + "trtllm::lora_grouped_gemm_cuda_graph", + [runner], + runner.tuning_config, + runner_inputs, + ) + return runner(runner_inputs, tactic=split_k) + def _forward_eager_mode( self, x: torch.Tensor, @@ -722,6 +787,198 @@ def _forward_eager_mode( return lora_output +class _LoraGroupedGemmRunner(TunableRunner): + """Tune split-K for one logical LoRA layer and token-count bucket.""" + + def __init__( + self, + layer: LoraLayer, + layer_idx: int, + input_hidden_size: int, + max_rank: int, + max_lora_size: int, + problem_count: int, + dtype: torch.dtype, + ): + self.layer = layer + self.layer_idx = layer_idx + self.input_hidden_size = input_hidden_size + self.max_rank = max_rank + self.max_lora_size = max_lora_size + self.problem_count = problem_count + self.dtype = dtype + self.layer_key = CudaGraphLoraParams.LoraLayerKey( + layer_idx=layer_idx, + module_ids=tuple(layer.lora_module_types), + ) + self.lora_params: Optional[Dict] = None + self.tuning_config = TuningConfig( + dynamic_tensor_specs=(DynamicTensorSpec( + 0, + 0, + get_last_power_of_2_num_tokens_buckets, + last_positive_power_of_2, + ), ), + inputs_pre_hook=self._prepare_synthetic_inputs, + ) + + def unique_id(self): + return ( + self.layer_idx, + tuple( + int(module_type) + for module_type in self.layer.lora_module_types), + tuple(self.layer.output_hidden_sizes), + self.input_hidden_size, + self.max_rank, + self.max_lora_size, + self.problem_count, + self.dtype, + ) + + def get_valid_tactics( + self, + inputs: List[torch.Tensor], + profile: OptimizationProfile, + **kwargs, + ) -> List[int]: + # input args are not needed to check valid tactics + del inputs, profile, kwargs + return list(_LORA_SPLIT_K_CANDIDATES) + + def copy_lora_params(self, lora_params: Dict) -> Dict: + """ + Copy the LoRA parameter hierarchy for this layer. + + Args: + lora_params: dict to be copied + + Returns: + Copied lora_params instance + """ + copied_lora_params = copy(lora_params) + cuda_graph_params = copy(lora_params['cuda_graph_params']) + layer_params = cuda_graph_params.get_layer_params(self.layer_key) + assert layer_params is not None + cuda_graph_params.layer_params = {self.layer_key: copy(layer_params)} + copied_lora_params['cuda_graph_params'] = cuda_graph_params + return copied_lora_params + + def _prepare_synthetic_inputs( + self, + inputs: List[torch.Tensor], + ) -> List[torch.Tensor]: + """ + Build one active-adapter problem for the requested token bucket. + + Args: + inputs: Input tensor + + Returns: + List of tensor input arguments for runner + + This method uses the local copy of lora_params in order to + create the list of tensor input arguments to be used by the + auto-tuner's forward. + """ + assert self.lora_params is not None + cuda_graph_params = self.lora_params['cuda_graph_params'] + layer_params = cuda_graph_params.get_layer_params(self.layer_key) + assert layer_params is not None + + token_carrier = inputs[0] + num_tokens = token_carrier.shape[0] + + b_ptrs = torch.zeros_like(layer_params.d_b_ptrs) + b_prime_ptrs = torch.zeros_like(layer_params.d_b_prime_ptrs) + keepalive = [] + for module_idx, output_size in enumerate( + self.layer.output_hidden_sizes): + lora_a = token_carrier.new_ones( + (self.max_rank, self.input_hidden_size)) + lora_b = token_carrier.new_ones((output_size, self.max_rank)) + b_ptrs[module_idx, 0] = lora_a.data_ptr() + b_prime_ptrs[module_idx, 0] = lora_b.data_ptr() + keepalive.extend((lora_a, lora_b)) + + slot_counts = torch.zeros_like(cuda_graph_params.slot_counts) + slot_counts[0] = num_tokens + slot_ranks = torch.zeros_like(cuda_graph_params.slot_ranks) + slot_ranks[0] = self.max_rank + slot_offsets_full = torch.zeros_like( + cuda_graph_params.slot_offsets_full) + slot_offsets_full[1:] = num_tokens + + return [ + token_carrier, + slot_counts, + slot_ranks, + slot_offsets_full, + b_ptrs, + b_prime_ptrs, + torch.arange(num_tokens, device=token_carrier.device), + layer_params.d_output_sizes, + layer_params.d_output_sizes_offset, + ] + keepalive + + def forward( + self, + /, + inputs: List[torch.Tensor], + *, + tactic: int = -1, + **kwargs, + ) -> torch.Tensor: + """ + Perform one auto-tuner LoraLayer forward pass. + + Args: + inputs: list of tensor input arguments + tactic: split-K value to be evaluated + + Returns: + LoRA output tensor + """ + del kwargs + assert self.lora_params is not None + lora_params = self.lora_params + x = inputs[0] + if len(inputs) > 1: + # Re-pack synthetic inputs into a local copy so that tactic + # evaluation does not modify the inference parameters. + lora_params = self.copy_lora_params(lora_params) + cuda_graph_params = lora_params['cuda_graph_params'] + layer_params = cuda_graph_params.get_layer_params(self.layer_key) + assert layer_params is not None + if self.dtype == torch.float8_e4m3fn: + cuda_graph_params.slot_ranks_host = ( + cuda_graph_params.slot_ranks_host.clone()) + cuda_graph_params.slot_ranks_host.zero_() + cuda_graph_params.slot_ranks_host[0] = self.max_rank + ( + x, + cuda_graph_params.slot_counts, + cuda_graph_params.slot_ranks, + cuda_graph_params.slot_offsets_full, + layer_params.d_b_ptrs, + layer_params.d_b_prime_ptrs, + cuda_graph_params.sorted_ids, + layer_params.d_output_sizes, + layer_params.d_output_sizes_offset, + *_keepalive, + ) = inputs + + split_k = _LORA_DEFAULT_SPLIT_K if tactic == -1 else tactic + output = self.layer._forward_cuda_graph_mode_impl( + x, + lora_params, + self.layer_idx, + split_k, + ) + assert isinstance(output, torch.Tensor) + return output + + class MoeLoraLayer(LoraLayer): """Marker LoraLayer for routed-expert MoE modules. diff --git a/tensorrt_llm/_torch/pyexecutor/model_engine.py b/tensorrt_llm/_torch/pyexecutor/model_engine.py index cbd37669a7fa..523d370f4bf7 100644 --- a/tensorrt_llm/_torch/pyexecutor/model_engine.py +++ b/tensorrt_llm/_torch/pyexecutor/model_engine.py @@ -1449,6 +1449,26 @@ def _get_full_general_warmup_requests( # Deduplicate the warmup_configs while keeping the order. return list(dict.fromkeys(warmup_configs)) + @contextmanager + def maybe_autotune_lora(self): + """Enable autotuning while warming up CUDA-graph LoRA kernels.""" + if not (self.llm_args.enable_autotuner + and self.cuda_graph_lora_manager is not None): + yield + return + + cache_path = os.environ.get("TLLM_AUTOTUNER_CACHE_PATH", None) + with autotune(cache_path=cache_path): + try: + yield + finally: + # Complete the PP cache hand-off even on ranks without a + # CUDA-graph-only tunable op. + autotuner = AutoTuner.get() + autotuner.cache_pp_recv() + autotuner.cache_pp_send() + autotuner.clean_pp_flag() + @with_warmup_flag @warmup_with_kv_cache_cleanup def warmup(self, resource_manager: ResourceManager) -> None: @@ -1549,7 +1569,8 @@ def warmup(self, resource_manager: ResourceManager) -> None: with self.cuda_graph_runner.allow_capture(): self.cuda_graph_runner.is_warmup_only = True try: - self._run_cuda_graph_warmup(resource_manager) + with self.maybe_autotune_lora(): + self._run_cuda_graph_warmup(resource_manager) finally: self.cuda_graph_runner.is_warmup_only = False self.cuda_graph_runner.padding_dummy_requests = {} diff --git a/tests/unittest/_torch/peft/test_lora_autotuner.py b/tests/unittest/_torch/peft/test_lora_autotuner.py new file mode 100644 index 000000000000..bcd661d3c38d --- /dev/null +++ b/tests/unittest/_torch/peft/test_lora_autotuner.py @@ -0,0 +1,272 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +from types import SimpleNamespace + +import pytest +import torch + +from tensorrt_llm._torch.autotuner import AutoTuner +from tensorrt_llm._torch.peft.lora import layer as lora_layer + + +def _make_runner( + layer_idx: int = 3, + input_hidden_size: int = 256, + dtype: torch.dtype = torch.float16, +) -> lora_layer._LoraGroupedGemmRunner: + layer = lora_layer.LoraLayer( + [ + lora_layer.LoraModuleType.ATTENTION_Q, + lora_layer.LoraModuleType.ATTENTION_K, + ], + [128, 64], + ) + return lora_layer._LoraGroupedGemmRunner( + layer=layer, + layer_idx=layer_idx, + input_hidden_size=input_hidden_size, + max_rank=16, + max_lora_size=4, + problem_count=8, + dtype=dtype, + ) + + +def test_lora_split_k_runner_uses_token_buckets(): + runner = _make_runner() + spec = runner.tuning_config.dynamic_tensor_specs[0] + + AutoTuner._find_nearest_profile.cache_clear() + profile = AutoTuner._find_nearest_profile( + (torch.Size((7, runner.input_hidden_size)),), + runner.tuning_config.dynamic_tensor_specs, + runner.tuning_config.constraint_specs, + runner.tuning_config.tune_max_num_tokens, + ) + + assert spec.input_idx == 0 + assert spec.dim_idx == 0 + assert profile[0][0] == 4 + + +def test_lora_split_k_runner_uses_default_tactic(monkeypatch): + runner = _make_runner() + runner.lora_params = {} + split_ks = [] + + def fake_forward_impl(x, lora_params, layer_idx, split_k): + del lora_params, layer_idx + split_ks.append(split_k) + return x + + monkeypatch.setattr(runner.layer, "_forward_cuda_graph_mode_impl", fake_forward_impl) + + x = torch.empty(2, runner.input_hidden_size) + assert runner([x], tactic=-1) is x + assert split_ks == [lora_layer._LORA_DEFAULT_SPLIT_K] + + +def test_fp8_lora_tuning_uses_private_synthetic_host_ranks(monkeypatch): + runner = _make_runner(dtype=torch.float8_e4m3fn) + layer_params = SimpleNamespace( + d_b_ptrs=torch.zeros((2, 4), dtype=torch.int64), + d_b_prime_ptrs=torch.zeros((2, 4), dtype=torch.int64), + d_output_sizes=torch.tensor([128, 64]), + d_output_sizes_offset=torch.tensor([0, 128]), + ) + cuda_graph_params = SimpleNamespace( + layer_params={runner.layer_key: layer_params}, + slot_ranks_host=torch.zeros(4, dtype=torch.int32), + ) + cuda_graph_params.get_layer_params = cuda_graph_params.layer_params.get + live_slot_ranks_host = cuda_graph_params.slot_ranks_host + runner.lora_params = runner.copy_lora_params({"cuda_graph_params": cuda_graph_params}) + + observed_slot_ranks_host = [] + + def fake_forward_impl(x, lora_params, layer_idx, split_k): + del layer_idx, split_k + observed_slot_ranks_host.append(lora_params["cuda_graph_params"].slot_ranks_host.clone()) + return x + + monkeypatch.setattr(runner.layer, "_forward_cuda_graph_mode_impl", fake_forward_impl) + + x = torch.empty((2, runner.input_hidden_size), dtype=torch.float8_e4m3fn) + synthetic_inputs = [ + x, + torch.tensor([2, 0, 0, 0]), + torch.tensor([runner.max_rank, 0, 0, 0]), + torch.tensor([0, 2, 2, 2, 2]), + layer_params.d_b_ptrs, + layer_params.d_b_prime_ptrs, + torch.arange(2), + layer_params.d_output_sizes, + layer_params.d_output_sizes_offset, + ] + + assert runner(synthetic_inputs, tactic=1) is x + torch.testing.assert_close( + observed_slot_ranks_host[0], + torch.tensor([runner.max_rank, 0, 0, 0], dtype=torch.int32), + ) + assert cuda_graph_params.slot_ranks_host is live_slot_ranks_host + torch.testing.assert_close( + live_slot_ranks_host, + torch.zeros(4, dtype=torch.int32), + ) + + +def test_lora_layer_reuses_runner_across_cuda_graph_warmups(monkeypatch): + layer = lora_layer.LoraLayer( + [lora_layer.LoraModuleType.ATTENTION_Q], + [128], + ) + layer_idx = 3 + layer_key = lora_layer.CudaGraphLoraParams.LoraLayerKey( + layer_idx=layer_idx, + module_ids=tuple(layer.lora_module_types), + ) + + class FakeLayerParams: + def __init__(self, max_lora_size: int = 4): + self.d_b_ptrs = torch.ones((1, max_lora_size), dtype=torch.int64) + self.d_b_prime_ptrs = torch.full((1, max_lora_size), 2, dtype=torch.int64) + self.d_output_sizes = torch.tensor([128]) + self.d_output_sizes_offset = torch.tensor([0]) + self.h_output_sizes = torch.tensor([128]) + + class FakeCudaGraphParams: + def __init__(self): + self.max_rank = 16 + self.max_lora_size = 4 + self.layer_info = {layer_key: object()} + self.layer_params = {layer_key: FakeLayerParams(self.max_lora_size)} + self.slot_counts = torch.tensor([4, 0, 0, 0]) + self.slot_ranks = torch.tensor([16, 0, 0, 0]) + self.slot_offsets_full = torch.tensor([0, 4, 4, 4, 4]) + self.sorted_ids = torch.arange(8) + + def get_layer_params(self, key): + assert key == layer_key + return self.layer_params.get(key) + + def get_problem_count(self, key): + assert key == layer_key + return 4 + + tuned_runners = [] + + class FakeTuner: + def choose_one(self, custom_op, runners, tuning_config, inputs): + assert custom_op == "trtllm::lora_grouped_gemm_cuda_graph" + assert tuning_config is runners[0].tuning_config + assert len(inputs) == 1 + tuned_runners.append(runners[0]) + return runners[0], 1 + + monkeypatch.setattr(lora_layer.AutoTuner, "get", staticmethod(lambda: FakeTuner())) + + parameter_fill_calls = [] + + def fake_parameter_fill(*args): + parameter_fill_calls.append(args) + + monkeypatch.setattr( + torch.ops.trtllm, + "lora_group_gemm_param_fill_row_reorder_fusion", + fake_parameter_fill, + ) + operator_split_ks = [] + + def fake_grouped_gemm(*args): + operator_split_ks.append(args[-1]) + + monkeypatch.setattr(torch.ops.trtllm, "lora_grouped_gemm_cuda_graph", fake_grouped_gemm) + + cuda_graph_params = FakeCudaGraphParams() + original_layer_params = cuda_graph_params.get_layer_params(layer_key) + original_slot_counts = cuda_graph_params.slot_counts + original_b_ptrs = original_layer_params.d_b_ptrs + lora_params = {"cuda_graph_params": cuda_graph_params} + warmup_lora_params = [] + for batch_size in (4, 8): + x = torch.empty(batch_size, 256) + output = layer._forward_cuda_graph_mode(x, lora_params, layer_idx) + assert output.shape == (batch_size, 128) + warmup_lora_params.append(layer._split_k_runner.lora_params) + + runner = layer._split_k_runner + assert runner is not None + assert tuned_runners == [runner, runner] + assert operator_split_ks == [1, 1] + assert warmup_lora_params[0] is not warmup_lora_params[1] + + runner_lora_params = warmup_lora_params[-1] + assert runner_lora_params is not None + runner_cuda_graph_params = runner_lora_params["cuda_graph_params"] + runner_layer_params = runner_cuda_graph_params.get_layer_params(layer_key) + assert runner_lora_params is not lora_params + assert runner_cuda_graph_params is not cuda_graph_params + assert runner_cuda_graph_params.layer_params is not cuda_graph_params.layer_params + assert runner_layer_params is not original_layer_params + + synthetic_inputs = runner._prepare_synthetic_inputs([torch.empty(2, 256)]) + runner(synthetic_inputs, tactic=2) + assert operator_split_ks == [1, 1, 2] + + synthetic_fill_args = parameter_fill_calls[-1] + assert synthetic_fill_args[16] is synthetic_inputs[1] + assert synthetic_fill_args[21] is synthetic_inputs[4] + assert runner_cuda_graph_params.slot_counts is original_slot_counts + assert runner_layer_params.d_b_ptrs is original_b_ptrs + assert cuda_graph_params.slot_counts is original_slot_counts + assert original_layer_params.d_b_ptrs is original_b_ptrs + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +def test_lora_autotuner_hook_builds_single_active_slot(): + runner = _make_runner() + num_tokens = 8 + carrier = torch.randn( + num_tokens, + runner.input_hidden_size, + dtype=runner.dtype, + device="cuda", + ) + layer_params = SimpleNamespace( + d_b_ptrs=torch.zeros((2, 4), dtype=torch.int64, device="cuda"), + d_b_prime_ptrs=torch.zeros((2, 4), dtype=torch.int64, device="cuda"), + d_output_sizes=torch.tensor([128, 64], dtype=torch.int32, device="cuda"), + d_output_sizes_offset=torch.tensor([0, 128], dtype=torch.int64, device="cuda"), + ) + cuda_graph_params = SimpleNamespace( + slot_counts=torch.zeros(4, dtype=torch.int32, device="cuda"), + slot_ranks=torch.zeros(4, dtype=torch.int32, device="cuda"), + slot_offsets_full=torch.zeros(5, dtype=torch.int64, device="cuda"), + layer_params={runner.layer_key: layer_params}, + ) + cuda_graph_params.get_layer_params = cuda_graph_params.layer_params.get + runner.lora_params = {"cuda_graph_params": cuda_graph_params} + + inputs = runner._prepare_synthetic_inputs([carrier]) + slot_counts, slot_ranks = inputs[1], inputs[2] + slot_offsets_full = inputs[3] + b_ptrs, b_prime_ptrs = inputs[4], inputs[5] + sorted_ids, output_hidden_sizes = inputs[6], inputs[7] + + assert slot_counts.tolist() == [num_tokens, 0, 0, 0] + assert slot_ranks.tolist() == [runner.max_rank, 0, 0, 0] + assert slot_offsets_full.tolist() == [ + 0, + num_tokens, + num_tokens, + num_tokens, + num_tokens, + ] + assert sorted_ids.tolist() == list(range(num_tokens)) + assert output_hidden_sizes.tolist() == [128, 64] + assert torch.all(b_ptrs[:, 0] != 0) + assert torch.all(b_prime_ptrs[:, 0] != 0) + assert torch.all(b_ptrs[:, 1:] == 0) + assert torch.all(b_prime_ptrs[:, 1:] == 0)