From 7215eff043b783a0a43b09c810c2b7c02e5addaf Mon Sep 17 00:00:00 2001 From: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> Date: Fri, 24 Jul 2026 08:12:03 -0700 Subject: [PATCH 01/10] [WiP] Auto-tuner support for LoRA splitk parameter Signed-off-by: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> --- tensorrt_llm/_torch/peft/lora/layer.py | 345 ++++++++++++++---- .../_torch/pyexecutor/model_engine.py | 17 +- .../_torch/peft/test_lora_autotuner.py | 94 +++++ 3 files changed, 386 insertions(+), 70 deletions(-) create mode 100644 tests/unittest/_torch/peft/test_lora_autotuner.py diff --git a/tensorrt_llm/_torch/peft/lora/layer.py b/tensorrt_llm/_torch/peft/lora/layer.py index f26ea410b535..8a150c7a49b9 100644 --- a/tensorrt_llm/_torch/peft/lora/layer.py +++ b/tensorrt_llm/_torch/peft/lora/layer.py @@ -16,12 +16,17 @@ from collections.abc import Callable from dataclasses import dataclass from enum import IntEnum -from typing import Dict, List, Optional +from os import getenv +from typing import Any, Dict, List, Optional, Tuple 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,11 @@ def _validate_fp8_lora_cuda_graph_alignment(slot_ranks_host: torch.Tensor, return min(hidden_size, min_active_rank) +# TODO: Potentially move this fallback to LoraConfig. +TRTLLM_SPLITK_VAL = int(getenv("TRTLLM_SPLITK_VAL", "8")) +_LORA_SPLIT_K_CANDIDATES = (1, 2, 4, 8, 16) + + @dataclass class GroupedGemmParamsOutput: in_sizes: Optional[torch.Tensor] = None @@ -232,6 +242,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_runners: Dict[Tuple, "_LoraGroupedGemmRunner"] = {} @staticmethod def forward_with_base( @@ -523,14 +534,9 @@ def _prepare_grouped_gemm_buffers_fused(self, splitk_offsets=splitk_offsets, reordered_input=reordered_input) - def _prepare_max_sizes_cpu(self, - cuda_graph_lora_params: CudaGraphLoraParams, - layer_key: CudaGraphLoraParams.LoraLayerKey, - bs: int, input_hidden_size: int): - layer_params = cuda_graph_lora_params.get_layer_params(layer_key) - shape_2d = (len(self.lora_module_types), - cuda_graph_lora_params.max_lora_size - ) # [num_layer_modules, max_lora_size] + def _prepare_max_sizes_cpu(self, bs: int, input_hidden_size: int, + max_lora_size: int, max_rank: int): + shape_2d = (len(self.lora_module_types), max_lora_size) shape_3d = shape_2d + (3, ) # dummy max sizes, on CPU host_max_in_sizes = torch.empty( @@ -540,15 +546,96 @@ def _prepare_max_sizes_cpu(self, host_max_in_sizes ) # m: batch_size, n: max_output_hidden_size, k: max_lora_rank host_max_in_sizes[:, :, 0] = bs - host_max_in_sizes[:, :, 1] = cuda_graph_lora_params.max_rank + host_max_in_sizes[:, :, 1] = max_rank host_max_in_sizes[:, :, 2] = input_hidden_size host_max_out_sizes[:, :, 0] = bs - host_max_out_sizes[:, :, 1] = layer_params.h_output_sizes.unsqueeze(1) - host_max_out_sizes[:, :, 2] = cuda_graph_lora_params.max_rank + host_max_out_sizes[:, :, 1] = torch.tensor( + self.output_hidden_sizes, + dtype=CudaGraphLoraParams.SIZES_DTYPE).unsqueeze(1) + host_max_out_sizes[:, :, 2] = max_rank return host_max_in_sizes, host_max_out_sizes + def _forward_cuda_graph_mode_impl( + self, + inputs: List[torch.Tensor], + max_lora_size: int, + max_rank: int, + problem_count: int, + min_kn: int, + split_k: int, + ) -> torch.Tensor: + """Run the complete CUDA-graph LoRA path with a fixed split-K.""" + x = inputs[0] + batch_size, hidden_size = x.shape + output_buffer = torch.empty( + (batch_size, sum(self.output_hidden_sizes)), + dtype=x.dtype, + device=x.device, + ) + params_input = GroupedGemmParamsInput( + x=x, + output_buffer=output_buffer, + intermediate_buffer=torch.empty( + (len(self.lora_module_types), batch_size, max_rank), + dtype=x.dtype, + device=x.device, + ), + max_lora_size=max_lora_size, + max_rank=max_rank, + slot_counts=inputs[1], + slot_ranks=inputs[2], + slot_offsets_full=inputs[3], + b_ptrs=inputs[4], + b_prime_ptrs=inputs[5], + sorted_ids=inputs[6], + output_hidden_sizes=inputs[7], + output_sizes_offset=inputs[8], + ) + host_max_in_sizes, host_max_out_sizes = self._prepare_max_sizes_cpu( + batch_size, + hidden_size, + max_lora_size, + max_rank, + ) + grouped_gemm_params = self._prepare_grouped_gemm_buffers_fused( + params_input) + + torch.ops.trtllm.lora_grouped_gemm_cuda_graph( + grouped_gemm_params.in_sizes, + grouped_gemm_params.out_sizes, + grouped_gemm_params.a_offset, + params_input.b_ptrs, + grouped_gemm_params.d_offset, + params_input.b_prime_ptrs, + grouped_gemm_params.d_prime_offset, + problem_count, + grouped_gemm_params.lda, + grouped_gemm_params.ldb, + 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, + params_input.x.dtype, + min_kn, + split_k, + ) + + # PyTorch does not implement index_copy_ for FP8 tensors. + if output_buffer.dtype == torch.float8_e4m3fn: + output_buffer = output_buffer.to(torch.bfloat16) + + restored_output = torch.empty_like(output_buffer) + restored_output.index_copy_( + 0, + params_input.sorted_ids[:batch_size], + output_buffer, + ) + return restored_output + def _forward_cuda_graph_mode( self, x: torch.Tensor, @@ -582,10 +669,8 @@ def _forward_cuda_graph_mode( if layer_params is None: return None # Pass-through for layers without LoRA modules - batch_size, hidden_size = x.shape[0], x.shape[-1] - num_layer_modules = len(self.lora_module_types) + _, hidden_size = x.shape max_rank = cuda_graph_params.max_rank - total_output_size = sum(self.output_hidden_sizes) if x.dtype == torch.float8_e4m3fn: min_kn = _validate_fp8_lora_cuda_graph_alignment( cuda_graph_params.slot_ranks_host, hidden_size, @@ -595,60 +680,46 @@ def _forward_cuda_graph_mode( hidden_size, 8, max_rank ) # TODO: hardcode to 8 for now, for alignments in kernels, might have alignment error if rank is less than 8! - output_buffer = torch.empty(batch_size, - total_output_size, - dtype=x.dtype, - device=x.device) - - host_max_in_sizes, host_max_out_sizes = self._prepare_max_sizes_cpu( - cuda_graph_params, layer_key, batch_size, hidden_size) - - # Intermediate buffer: [num_layer_modules, batch_size, max_rank] - intermediate_buffer = torch.empty( - [num_layer_modules, batch_size, max_rank], - dtype=x.dtype, - device=x.device) - - params_fill_input = GroupedGemmParamsInput( - x=x, - output_buffer=output_buffer, - intermediate_buffer=intermediate_buffer, - max_lora_size=cuda_graph_params.max_lora_size, - max_rank=cuda_graph_params.max_rank, - slot_counts=cuda_graph_params.slot_counts, - slot_ranks=cuda_graph_params.slot_ranks, - slot_offsets_full=cuda_graph_params.slot_offsets_full, - b_ptrs=layer_params.d_b_ptrs, - b_prime_ptrs=layer_params.d_b_prime_ptrs, - sorted_ids=cuda_graph_params.sorted_ids, - output_hidden_sizes=layer_params.d_output_sizes, - output_sizes_offset=layer_params.d_output_sizes_offset) - grouped_gemm_params = self._prepare_grouped_gemm_buffers_fused( - params_fill_input) - - torch.ops.trtllm.lora_grouped_gemm_cuda_graph( - grouped_gemm_params.in_sizes, grouped_gemm_params.out_sizes, - grouped_gemm_params.a_offset, layer_params.d_b_ptrs, - grouped_gemm_params.d_offset, layer_params.d_b_prime_ptrs, - grouped_gemm_params.d_prime_offset, - cuda_graph_params.get_problem_count(layer_key), - grouped_gemm_params.lda, grouped_gemm_params.ldb, - 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) - - # PyTorch does not implement index_copy_ for FP8 tensors. - if output_buffer.dtype == torch.float8_e4m3fn: - output_buffer = output_buffer.to(torch.bfloat16) - - # TODO: move to kernel - # sorted_ids is a permutation, so index_copy_ initializes every row. - restored_output = torch.empty_like(output_buffer) - restored_output.index_copy_(0, - cuda_graph_params.sorted_ids[:batch_size], - output_buffer) - return restored_output + problem_count = cuda_graph_params.get_problem_count(layer_key) + runner_key = ( + layer_idx, + hidden_size, + max_rank, + cuda_graph_params.max_lora_size, + problem_count, + x.dtype, + min_kn, + ) + if runner_key not in self._split_k_runners: + self._split_k_runners[runner_key] = _LoraGroupedGemmRunner( + layer=self, + layer_idx=layer_idx, + input_hidden_size=hidden_size, + max_rank=max_rank, + max_lora_size=cuda_graph_params.max_lora_size, + problem_count=problem_count, + dtype=x.dtype, + min_kn=min_kn, + ) + runner = self._split_k_runners[runner_key] + runner_inputs = [ + 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, + ] + _, 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, @@ -722,6 +793,142 @@ 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, + min_kn: int, + ): + 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.min_kn = min_kn + 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, + self.min_kn, + ) + + def get_valid_tactics( + self, + inputs: List[torch.Tensor], + profile: OptimizationProfile, + **kwargs, + ) -> List[int]: + del inputs, profile, kwargs + k_tiles = max(1, self.input_hidden_size // 64) + return [ + split_k for split_k in _LORA_SPLIT_K_CANDIDATES + if split_k <= k_tiles + ] + + def _prepare_synthetic_inputs( + self, inputs: List[torch.Tensor]) -> List[torch.Tensor]: + """Build one active-adapter problem for the requested token bucket.""" + token_carrier = inputs[0] + num_tokens = token_carrier.shape[0] + device = token_carrier.device + module_count = len(self.layer.lora_module_types) + shape_2d = (module_count, self.max_lora_size) + + b_ptrs = torch.zeros(shape_2d, + dtype=CudaGraphLoraParams.PTR_DTYPE, + device=device) + b_prime_ptrs = torch.zeros_like(b_ptrs) + keepalive = [] + for module_idx, output_size in enumerate( + self.layer.output_hidden_sizes): + lora_a = torch.ones((self.max_rank, self.input_hidden_size), + dtype=self.dtype, + device=device) + lora_b = torch.ones((output_size, self.max_rank), + dtype=self.dtype, + device=device) + 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(self.max_lora_size, + dtype=CudaGraphLoraParams.SIZES_DTYPE, + device=device) + slot_counts[0] = num_tokens + slot_ranks = torch.zeros_like(slot_counts) + slot_ranks[0] = self.max_rank + slot_offsets_full = torch.zeros( + self.max_lora_size + 1, + dtype=CudaGraphLoraParams.PTR_DTYPE, + device=device) + slot_offsets_full[1:] = num_tokens + + output_hidden_sizes = torch.tensor( + self.layer.output_hidden_sizes, + dtype=CudaGraphLoraParams.SIZES_DTYPE, + device=device) + output_sizes_offset = CudaGraphLoraParams.get_offset_from_counts( + output_hidden_sizes).to(dtype=CudaGraphLoraParams.PTR_DTYPE) + + return [ + token_carrier, + slot_counts, + slot_ranks, + slot_offsets_full, + b_ptrs, + b_prime_ptrs, + torch.arange(num_tokens, dtype=torch.int64, device=device), + output_hidden_sizes, + output_sizes_offset, + ] + keepalive + + def forward( + self, + /, + inputs: List[torch.Tensor], + *, + tactic: int = -1, + **kwargs, + ) -> torch.Tensor: + del kwargs + split_k = TRTLLM_SPLITK_VAL if tactic == -1 else tactic + return self.layer._forward_cuda_graph_mode_impl( + inputs, + self.max_lora_size, + self.max_rank, + self.problem_count, + self.min_kn, + split_k, + ) + + 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..0124e4d70440 100644 --- a/tensorrt_llm/_torch/pyexecutor/model_engine.py +++ b/tensorrt_llm/_torch/pyexecutor/model_engine.py @@ -1549,7 +1549,22 @@ 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) + cache_path = os.environ.get("TLLM_AUTOTUNER_CACHE_PATH", + None) + lora_autotuner_enabled = ( + self.llm_args.enable_autotuner + and self.cuda_graph_lora_manager is not None) + autotune_ctx = (autotune(cache_path=cache_path) + if lora_autotuner_enabled else + contextlib.nullcontext()) + with autotune_ctx: + self._run_cuda_graph_warmup(resource_manager) + if lora_autotuner_enabled: + # Complete the PP cache hand-off even on ranks + # without a CUDA-graph-only tunable op. + AutoTuner.get().cache_pp_recv() + AutoTuner.get().cache_pp_send() + AutoTuner.get().clean_pp_flag() 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..470490c91d97 --- /dev/null +++ b/tests/unittest/_torch/peft/test_lora_autotuner.py @@ -0,0 +1,94 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +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, +) -> 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=torch.float16, + ) + + +def test_lora_split_k_runner_identity_is_layer_specific(): + runner = _make_runner(layer_idx=3) + other_layer_runner = _make_runner(layer_idx=4) + + assert runner.unique_id() != other_layer_runner.unique_id() + + +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_prunes_splits_larger_than_k_tiles(): + runner = _make_runner(input_hidden_size=256) + + assert runner.get_valid_tactics([], None) == [1, 2, 4] + + +@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", + ) + + 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) From 3140931e2c3c122b9032857f78409fad5c04f1be Mon Sep 17 00:00:00 2001 From: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> Date: Mon, 27 Jul 2026 06:13:57 -0700 Subject: [PATCH 02/10] Streamline LoRA autotuning logic in model_engine.py Signed-off-by: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> --- .../_torch/pyexecutor/model_engine.py | 36 +++++++++++-------- 1 file changed, 21 insertions(+), 15 deletions(-) diff --git a/tensorrt_llm/_torch/pyexecutor/model_engine.py b/tensorrt_llm/_torch/pyexecutor/model_engine.py index 0124e4d70440..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,22 +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: - cache_path = os.environ.get("TLLM_AUTOTUNER_CACHE_PATH", - None) - lora_autotuner_enabled = ( - self.llm_args.enable_autotuner - and self.cuda_graph_lora_manager is not None) - autotune_ctx = (autotune(cache_path=cache_path) - if lora_autotuner_enabled else - contextlib.nullcontext()) - with autotune_ctx: + with self.maybe_autotune_lora(): self._run_cuda_graph_warmup(resource_manager) - if lora_autotuner_enabled: - # Complete the PP cache hand-off even on ranks - # without a CUDA-graph-only tunable op. - AutoTuner.get().cache_pp_recv() - AutoTuner.get().cache_pp_send() - AutoTuner.get().clean_pp_flag() finally: self.cuda_graph_runner.is_warmup_only = False self.cuda_graph_runner.padding_dummy_requests = {} From 8e444eca73a7b3d31fda2d843f8781e90db2f7b6 Mon Sep 17 00:00:00 2001 From: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> Date: Thu, 30 Jul 2026 00:12:08 -0700 Subject: [PATCH 03/10] Use single autotuner runner instance per LoraLayer Signed-off-by: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> --- tensorrt_llm/_torch/peft/lora/layer.py | 32 ++++----- .../_torch/peft/test_lora_autotuner.py | 70 +++++++++++++++++++ 2 files changed, 82 insertions(+), 20 deletions(-) diff --git a/tensorrt_llm/_torch/peft/lora/layer.py b/tensorrt_llm/_torch/peft/lora/layer.py index 8a150c7a49b9..e18cf86d0b65 100644 --- a/tensorrt_llm/_torch/peft/lora/layer.py +++ b/tensorrt_llm/_torch/peft/lora/layer.py @@ -17,7 +17,7 @@ from dataclasses import dataclass from enum import IntEnum from os import getenv -from typing import Any, Dict, List, Optional, Tuple +from typing import Dict, List, Optional import torch @@ -242,7 +242,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_runners: Dict[Tuple, "_LoraGroupedGemmRunner"] = {} + self._split_k_runner: Optional["_LoraGroupedGemmRunner"] = None @staticmethod def forward_with_base( @@ -681,17 +681,8 @@ def _forward_cuda_graph_mode( ) # TODO: hardcode to 8 for now, for alignments in kernels, might have alignment error if rank is less than 8! problem_count = cuda_graph_params.get_problem_count(layer_key) - runner_key = ( - layer_idx, - hidden_size, - max_rank, - cuda_graph_params.max_lora_size, - problem_count, - x.dtype, - min_kn, - ) - if runner_key not in self._split_k_runners: - self._split_k_runners[runner_key] = _LoraGroupedGemmRunner( + if self._split_k_runner is None: + self._split_k_runner = _LoraGroupedGemmRunner( layer=self, layer_idx=layer_idx, input_hidden_size=hidden_size, @@ -701,7 +692,8 @@ def _forward_cuda_graph_mode( dtype=x.dtype, min_kn=min_kn, ) - runner = self._split_k_runners[runner_key] + runner = self._split_k_runner + runner.min_kn = min_kn runner_inputs = [ x, cuda_graph_params.slot_counts, @@ -828,8 +820,9 @@ def __init__( def unique_id(self): return ( self.layer_idx, - tuple(int(module_type) - for module_type in self.layer.lora_module_types), + 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, @@ -884,10 +877,9 @@ def _prepare_synthetic_inputs( slot_counts[0] = num_tokens slot_ranks = torch.zeros_like(slot_counts) slot_ranks[0] = self.max_rank - slot_offsets_full = torch.zeros( - self.max_lora_size + 1, - dtype=CudaGraphLoraParams.PTR_DTYPE, - device=device) + slot_offsets_full = torch.zeros(self.max_lora_size + 1, + dtype=CudaGraphLoraParams.PTR_DTYPE, + device=device) slot_offsets_full[1:] = num_tokens output_hidden_sizes = torch.tensor( diff --git a/tests/unittest/_torch/peft/test_lora_autotuner.py b/tests/unittest/_torch/peft/test_lora_autotuner.py index 470490c91d97..7b13e62b0951 100644 --- a/tests/unittest/_torch/peft/test_lora_autotuner.py +++ b/tests/unittest/_torch/peft/test_lora_autotuner.py @@ -60,6 +60,76 @@ def test_lora_split_k_runner_prunes_splits_larger_than_k_tiles(): assert runner.get_valid_tactics([], None) == [1, 2, 4] +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: + d_b_ptrs = None + d_b_prime_ptrs = None + d_output_sizes = None + d_output_sizes_offset = None + + class FakeCudaGraphParams: + layer_info = {layer_key: object()} + max_rank = 16 + max_lora_size = 4 + slot_counts = None + slot_ranks = None + slot_offsets_full = None + sorted_ids = None + + def get_layer_params(self, key): + assert key == layer_key + return FakeLayerParams() + + def get_problem_count(self, key): + assert key == layer_key + return 4 + + runners_created = [] + + class FakeRunner: + def __init__(self, **kwargs): + self.tuning_config = object() + self.kwargs = kwargs + self.calls = [] + runners_created.append(self) + + def __call__(self, inputs, *, tactic): + self.calls.append((inputs, tactic)) + return inputs[0] + + 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 + tuned_runners.append(runners[0]) + return runners[0], 1 + + monkeypatch.setattr(lora_layer, "_LoraGroupedGemmRunner", FakeRunner) + monkeypatch.setattr(lora_layer.AutoTuner, "get", staticmethod(lambda: FakeTuner())) + + lora_params = {"cuda_graph_params": FakeCudaGraphParams()} + for batch_size in (4, 8): + x = torch.empty(batch_size, 256) + assert layer._forward_cuda_graph_mode(x, lora_params, layer_idx) is x + + assert len(runners_created) == 1 + assert layer._split_k_runner is runners_created[0] + assert tuned_runners == [runners_created[0], runners_created[0]] + assert len(runners_created[0].calls) == 2 + + @pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") def test_lora_autotuner_hook_builds_single_active_slot(): runner = _make_runner() From 199722599917a51e08b1836a2dc3a472482e4a6a Mon Sep 17 00:00:00 2001 From: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> Date: Thu, 30 Jul 2026 00:47:43 -0700 Subject: [PATCH 04/10] Streamline input handling for LoRA autotuner Signed-off-by: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> --- tensorrt_llm/_torch/peft/lora/layer.py | 319 ++++++++++-------- .../_torch/peft/test_lora_autotuner.py | 101 ++++-- 2 files changed, 244 insertions(+), 176 deletions(-) diff --git a/tensorrt_llm/_torch/peft/lora/layer.py b/tensorrt_llm/_torch/peft/lora/layer.py index e18cf86d0b65..a47fdb918456 100644 --- a/tensorrt_llm/_torch/peft/lora/layer.py +++ b/tensorrt_llm/_torch/peft/lora/layer.py @@ -14,9 +14,9 @@ # limitations under the License. from collections.abc import Callable +from copy import copy from dataclasses import dataclass from enum import IntEnum -from os import getenv from typing import Dict, List, Optional import torch @@ -27,7 +27,7 @@ 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 +from .cuda_graph_lora_params import CudaGraphLoraParams, LoraLayerParams _FP8_LORA_TMA_ALIGNMENT = 16 @@ -69,8 +69,7 @@ def _validate_fp8_lora_cuda_graph_alignment(slot_ranks_host: torch.Tensor, return min(hidden_size, min_active_rank) -# TODO: Potentially move this fallback to LoraConfig. -TRTLLM_SPLITK_VAL = int(getenv("TRTLLM_SPLITK_VAL", "8")) +_LORA_DEFAULT_SPLIT_K = 16 _LORA_SPLIT_K_CANDIDATES = (1, 2, 4, 8, 16) @@ -534,9 +533,14 @@ def _prepare_grouped_gemm_buffers_fused(self, splitk_offsets=splitk_offsets, reordered_input=reordered_input) - def _prepare_max_sizes_cpu(self, bs: int, input_hidden_size: int, - max_lora_size: int, max_rank: int): - shape_2d = (len(self.lora_module_types), max_lora_size) + def _prepare_max_sizes_cpu(self, + cuda_graph_lora_params: CudaGraphLoraParams, + layer_key: CudaGraphLoraParams.LoraLayerKey, + bs: int, input_hidden_size: int): + layer_params = cuda_graph_lora_params.get_layer_params(layer_key) + shape_2d = (len(self.lora_module_types), + cuda_graph_lora_params.max_lora_size + ) # [num_layer_modules, max_lora_size] shape_3d = shape_2d + (3, ) # dummy max sizes, on CPU host_max_in_sizes = torch.empty( @@ -546,94 +550,104 @@ def _prepare_max_sizes_cpu(self, bs: int, input_hidden_size: int, host_max_in_sizes ) # m: batch_size, n: max_output_hidden_size, k: max_lora_rank host_max_in_sizes[:, :, 0] = bs - host_max_in_sizes[:, :, 1] = max_rank + host_max_in_sizes[:, :, 1] = cuda_graph_lora_params.max_rank host_max_in_sizes[:, :, 2] = input_hidden_size host_max_out_sizes[:, :, 0] = bs - host_max_out_sizes[:, :, 1] = torch.tensor( - self.output_hidden_sizes, - dtype=CudaGraphLoraParams.SIZES_DTYPE).unsqueeze(1) - host_max_out_sizes[:, :, 2] = max_rank + host_max_out_sizes[:, :, 1] = layer_params.h_output_sizes.unsqueeze(1) + host_max_out_sizes[:, :, 2] = cuda_graph_lora_params.max_rank return host_max_in_sizes, host_max_out_sizes def _forward_cuda_graph_mode_impl( self, - inputs: List[torch.Tensor], - max_lora_size: int, - max_rank: int, - problem_count: int, - min_kn: int, + x: torch.Tensor, + lora_params: Dict, + layer_idx: int, split_k: int, - ) -> torch.Tensor: + ) -> Optional[torch.Tensor]: """Run the complete CUDA-graph LoRA path with a fixed split-K.""" - x = inputs[0] - batch_size, hidden_size = x.shape - output_buffer = torch.empty( - (batch_size, sum(self.output_hidden_sizes)), + 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 + + layer_params = cuda_graph_params.get_layer_params(layer_key) + + # Skip layers that don't have LoRA modules + if layer_params is None: + return None # Pass-through for layers without LoRA modules + + batch_size, hidden_size = x.shape[0], x.shape[-1] + num_layer_modules = len(self.lora_module_types) + max_rank = cuda_graph_params.max_rank + total_output_size = sum(self.output_hidden_sizes) + if x.dtype == torch.float8_e4m3fn: + min_kn = _validate_fp8_lora_cuda_graph_alignment( + cuda_graph_params.slot_ranks_host, hidden_size, + self.output_hidden_sizes, max_rank) + else: + min_kn = min( + hidden_size, 8, max_rank + ) # TODO: hardcode to 8 for now, for alignments in kernels, might have alignment error if rank is less than 8! + + output_buffer = torch.empty(batch_size, + total_output_size, + dtype=x.dtype, + device=x.device) + + host_max_in_sizes, host_max_out_sizes = self._prepare_max_sizes_cpu( + cuda_graph_params, layer_key, batch_size, hidden_size) + + # Intermediate buffer: [num_layer_modules, batch_size, max_rank] + intermediate_buffer = torch.empty( + [num_layer_modules, batch_size, max_rank], dtype=x.dtype, - device=x.device, - ) - params_input = GroupedGemmParamsInput( + device=x.device) + + params_fill_input = GroupedGemmParamsInput( x=x, output_buffer=output_buffer, - intermediate_buffer=torch.empty( - (len(self.lora_module_types), batch_size, max_rank), - dtype=x.dtype, - device=x.device, - ), - max_lora_size=max_lora_size, - max_rank=max_rank, - slot_counts=inputs[1], - slot_ranks=inputs[2], - slot_offsets_full=inputs[3], - b_ptrs=inputs[4], - b_prime_ptrs=inputs[5], - sorted_ids=inputs[6], - output_hidden_sizes=inputs[7], - output_sizes_offset=inputs[8], - ) - host_max_in_sizes, host_max_out_sizes = self._prepare_max_sizes_cpu( - batch_size, - hidden_size, - max_lora_size, - max_rank, - ) + intermediate_buffer=intermediate_buffer, + max_lora_size=cuda_graph_params.max_lora_size, + max_rank=cuda_graph_params.max_rank, + slot_counts=cuda_graph_params.slot_counts, + slot_ranks=cuda_graph_params.slot_ranks, + slot_offsets_full=cuda_graph_params.slot_offsets_full, + b_ptrs=layer_params.d_b_ptrs, + b_prime_ptrs=layer_params.d_b_prime_ptrs, + sorted_ids=cuda_graph_params.sorted_ids, + output_hidden_sizes=layer_params.d_output_sizes, + output_sizes_offset=layer_params.d_output_sizes_offset) grouped_gemm_params = self._prepare_grouped_gemm_buffers_fused( - params_input) + params_fill_input) torch.ops.trtllm.lora_grouped_gemm_cuda_graph( - grouped_gemm_params.in_sizes, - grouped_gemm_params.out_sizes, - grouped_gemm_params.a_offset, - params_input.b_ptrs, - grouped_gemm_params.d_offset, - params_input.b_prime_ptrs, + grouped_gemm_params.in_sizes, grouped_gemm_params.out_sizes, + grouped_gemm_params.a_offset, layer_params.d_b_ptrs, + grouped_gemm_params.d_offset, layer_params.d_b_prime_ptrs, grouped_gemm_params.d_prime_offset, - problem_count, - grouped_gemm_params.lda, - grouped_gemm_params.ldb, - 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, - params_input.x.dtype, - min_kn, - split_k, - ) + cuda_graph_params.get_problem_count(layer_key), + grouped_gemm_params.lda, grouped_gemm_params.ldb, + 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, split_k) # PyTorch does not implement index_copy_ for FP8 tensors. if output_buffer.dtype == torch.float8_e4m3fn: output_buffer = output_buffer.to(torch.bfloat16) + # TODO: move to kernel + # sorted_ids is a permutation, so index_copy_ initializes every row. restored_output = torch.empty_like(output_buffer) - restored_output.index_copy_( - 0, - params_input.sorted_ids[:batch_size], - output_buffer, - ) + restored_output.index_copy_(0, + cuda_graph_params.sorted_ids[:batch_size], + output_buffer) return restored_output def _forward_cuda_graph_mode( @@ -663,48 +677,23 @@ def _forward_cuda_graph_mode( if not cuda_graph_params or not cuda_graph_params.layer_info or layer_key not in cuda_graph_params.layer_info: return None - layer_params = cuda_graph_params.get_layer_params(layer_key) - # 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 - - _, hidden_size = x.shape - max_rank = cuda_graph_params.max_rank - if x.dtype == torch.float8_e4m3fn: - min_kn = _validate_fp8_lora_cuda_graph_alignment( - cuda_graph_params.slot_ranks_host, hidden_size, - self.output_hidden_sizes, max_rank) - else: - min_kn = min( - hidden_size, 8, max_rank - ) # TODO: hardcode to 8 for now, for alignments in kernels, might have alignment error if rank is less than 8! - - problem_count = cuda_graph_params.get_problem_count(layer_key) if self._split_k_runner is None: self._split_k_runner = _LoraGroupedGemmRunner( layer=self, layer_idx=layer_idx, - input_hidden_size=hidden_size, - max_rank=max_rank, + input_hidden_size=x.shape[1], + max_rank=cuda_graph_params.max_rank, max_lora_size=cuda_graph_params.max_lora_size, - problem_count=problem_count, + problem_count=cuda_graph_params.get_problem_count(layer_key), dtype=x.dtype, - min_kn=min_kn, ) + runner = self._split_k_runner - runner.min_kn = min_kn - runner_inputs = [ - 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, - ] + runner_inputs = runner.prepare_inputs(x, lora_params, cuda_graph_params) _, split_k = AutoTuner.get().choose_one( "trtllm::lora_grouped_gemm_cuda_graph", [runner], @@ -797,7 +786,6 @@ def __init__( max_lora_size: int, problem_count: int, dtype: torch.dtype, - min_kn: int, ): self.layer = layer self.layer_idx = layer_idx @@ -806,7 +794,13 @@ def __init__( self.max_lora_size = max_lora_size self.problem_count = problem_count self.dtype = dtype - self.min_kn = min_kn + self.layer_key = CudaGraphLoraParams.LoraLayerKey( + layer_idx=layer_idx, + module_ids=tuple(layer.lora_module_types), + ) + self.lora_params: Optional[Dict] = None + self.cuda_graph_params: Optional[CudaGraphLoraParams] = None + self.layer_params: Optional[LoraLayerParams] = None self.tuning_config = TuningConfig( dynamic_tensor_specs=(DynamicTensorSpec( 0, @@ -829,7 +823,6 @@ def unique_id(self): self.max_lora_size, self.problem_count, self.dtype, - self.min_kn, ) def get_valid_tactics( @@ -845,50 +838,57 @@ def get_valid_tactics( if split_k <= k_tiles ] + def prepare_inputs( + self, + x: torch.Tensor, + lora_params: Dict, + cuda_graph_params: CudaGraphLoraParams, + ) -> List[torch.Tensor]: + """Copy the LoRA parameters and pack the autotuner tensor inputs.""" + self.lora_params = copy(lora_params) + self.cuda_graph_params = copy(cuda_graph_params) + + layer_params = cuda_graph_params.get_layer_params(self.layer_key) + assert layer_params is not None + self.layer_params = copy(layer_params) + self.cuda_graph_params.layer_params = { + self.layer_key: self.layer_params + } + self.lora_params['cuda_graph_params'] = self.cuda_graph_params + + return [x] + def _prepare_synthetic_inputs( - self, inputs: List[torch.Tensor]) -> List[torch.Tensor]: + self, + inputs: List[torch.Tensor], + ) -> List[torch.Tensor]: """Build one active-adapter problem for the requested token bucket.""" + assert self.cuda_graph_params is not None + assert self.layer_params is not None + token_carrier = inputs[0] num_tokens = token_carrier.shape[0] - device = token_carrier.device - module_count = len(self.layer.lora_module_types) - shape_2d = (module_count, self.max_lora_size) - - b_ptrs = torch.zeros(shape_2d, - dtype=CudaGraphLoraParams.PTR_DTYPE, - device=device) - b_prime_ptrs = torch.zeros_like(b_ptrs) + + b_ptrs = torch.zeros_like(self.layer_params.d_b_ptrs) + b_prime_ptrs = torch.zeros_like(self.layer_params.d_b_prime_ptrs) keepalive = [] for module_idx, output_size in enumerate( self.layer.output_hidden_sizes): - lora_a = torch.ones((self.max_rank, self.input_hidden_size), - dtype=self.dtype, - device=device) - lora_b = torch.ones((output_size, self.max_rank), - dtype=self.dtype, - device=device) + 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(self.max_lora_size, - dtype=CudaGraphLoraParams.SIZES_DTYPE, - device=device) + slot_counts = torch.zeros_like(self.cuda_graph_params.slot_counts) slot_counts[0] = num_tokens - slot_ranks = torch.zeros_like(slot_counts) + slot_ranks = torch.zeros_like(self.cuda_graph_params.slot_ranks) slot_ranks[0] = self.max_rank - slot_offsets_full = torch.zeros(self.max_lora_size + 1, - dtype=CudaGraphLoraParams.PTR_DTYPE, - device=device) + slot_offsets_full = torch.zeros_like( + self.cuda_graph_params.slot_offsets_full) slot_offsets_full[1:] = num_tokens - output_hidden_sizes = torch.tensor( - self.layer.output_hidden_sizes, - dtype=CudaGraphLoraParams.SIZES_DTYPE, - device=device) - output_sizes_offset = CudaGraphLoraParams.get_offset_from_counts( - output_hidden_sizes).to(dtype=CudaGraphLoraParams.PTR_DTYPE) - return [ token_carrier, slot_counts, @@ -896,9 +896,9 @@ def _prepare_synthetic_inputs( slot_offsets_full, b_ptrs, b_prime_ptrs, - torch.arange(num_tokens, dtype=torch.int64, device=device), - output_hidden_sizes, - output_sizes_offset, + torch.arange(num_tokens, device=token_carrier.device), + self.layer_params.d_output_sizes, + self.layer_params.d_output_sizes_offset, ] + keepalive def forward( @@ -910,15 +910,44 @@ def forward( **kwargs, ) -> torch.Tensor: del kwargs - split_k = TRTLLM_SPLITK_VAL if tactic == -1 else tactic - return self.layer._forward_cuda_graph_mode_impl( - inputs, - self.max_lora_size, - self.max_rank, - self.problem_count, - self.min_kn, + assert self.lora_params is not None + assert self.cuda_graph_params is not None + assert self.layer_params is not None + + lora_params = copy(self.lora_params) + cuda_graph_params = copy(self.cuda_graph_params) + layer_params = copy(self.layer_params) + cuda_graph_params.layer_params = {self.layer_key: layer_params} + lora_params['cuda_graph_params'] = cuda_graph_params + + x = inputs[0] + if len(inputs) > 1: + 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): diff --git a/tests/unittest/_torch/peft/test_lora_autotuner.py b/tests/unittest/_torch/peft/test_lora_autotuner.py index 7b13e62b0951..5be2f6a0d8e9 100644 --- a/tests/unittest/_torch/peft/test_lora_autotuner.py +++ b/tests/unittest/_torch/peft/test_lora_autotuner.py @@ -1,6 +1,8 @@ # 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 @@ -72,62 +74,88 @@ def test_lora_layer_reuses_runner_across_cuda_graph_warmups(monkeypatch): ) class FakeLayerParams: - d_b_ptrs = None - d_b_prime_ptrs = None - d_output_sizes = None - d_output_sizes_offset = None + def __init__(self): + self.d_b_ptrs = torch.tensor([1]) + self.d_b_prime_ptrs = torch.tensor([2]) + self.d_output_sizes = torch.tensor([128]) + self.d_output_sizes_offset = torch.tensor([0]) + self.h_output_sizes = torch.tensor([128]) class FakeCudaGraphParams: - layer_info = {layer_key: object()} - max_rank = 16 - max_lora_size = 4 - slot_counts = None - slot_ranks = None - slot_offsets_full = None - sorted_ids = None + def __init__(self): + self.layer_info = {layer_key: object()} + self.layer_params = {layer_key: FakeLayerParams()} + self.max_rank = 16 + self.max_lora_size = 4 + 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 FakeLayerParams() + return self.layer_params.get(key) def get_problem_count(self, key): assert key == layer_key return 4 - runners_created = [] - - class FakeRunner: - def __init__(self, **kwargs): - self.tuning_config = object() - self.kwargs = kwargs - self.calls = [] - runners_created.append(self) - - def __call__(self, inputs, *, tactic): - self.calls.append((inputs, tactic)) - return inputs[0] - 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, "_LoraGroupedGemmRunner", FakeRunner) monkeypatch.setattr(lora_layer.AutoTuner, "get", staticmethod(lambda: FakeTuner())) - lora_params = {"cuda_graph_params": FakeCudaGraphParams()} + forwarded_params = [] + + def fake_forward_impl(x, runner_lora_params, forwarded_layer_idx, split_k): + forwarded_params.append((runner_lora_params, split_k)) + assert forwarded_layer_idx == layer_idx + return x + + monkeypatch.setattr(layer, "_forward_cuda_graph_mode_impl", fake_forward_impl) + + 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} for batch_size in (4, 8): x = torch.empty(batch_size, 256) assert layer._forward_cuda_graph_mode(x, lora_params, layer_idx) is x - assert len(runners_created) == 1 - assert layer._split_k_runner is runners_created[0] - assert tuned_runners == [runners_created[0], runners_created[0]] - assert len(runners_created[0].calls) == 2 + runner = layer._split_k_runner + assert runner is not None + assert tuned_runners == [runner, runner] + assert len(forwarded_params) == 2 + assert forwarded_params[0][0] is not forwarded_params[1][0] + + runner_lora_params = forwarded_params[-1][0] + 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) + + synthetic_lora_params = forwarded_params[-1][0] + synthetic_cuda_graph_params = synthetic_lora_params["cuda_graph_params"] + synthetic_layer_params = synthetic_cuda_graph_params.get_layer_params(layer_key) + assert synthetic_cuda_graph_params.slot_counts is synthetic_inputs[1] + assert synthetic_layer_params.d_b_ptrs 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") @@ -140,6 +168,17 @@ def test_lora_autotuner_hook_builds_single_active_slot(): dtype=runner.dtype, device="cuda", ) + runner.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"), + ) + runner.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"), + ) inputs = runner._prepare_synthetic_inputs([carrier]) slot_counts, slot_ranks = inputs[1], inputs[2] From 0b879341d9f4a15aebdb3e95707d572c7a8f8045 Mon Sep 17 00:00:00 2001 From: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> Date: Fri, 31 Jul 2026 02:23:09 -0700 Subject: [PATCH 05/10] Add docstrings and simplify valid tactic selection Signed-off-by: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> --- tensorrt_llm/_torch/peft/lora/layer.py | 64 ++++++++++++++++++++++---- 1 file changed, 56 insertions(+), 8 deletions(-) diff --git a/tensorrt_llm/_torch/peft/lora/layer.py b/tensorrt_llm/_torch/peft/lora/layer.py index a47fdb918456..396dac57eaa5 100644 --- a/tensorrt_llm/_torch/peft/lora/layer.py +++ b/tensorrt_llm/_torch/peft/lora/layer.py @@ -566,7 +566,19 @@ def _forward_cuda_graph_mode_impl( layer_idx: int, split_k: int, ) -> Optional[torch.Tensor]: - """Run the complete CUDA-graph LoRA path with a fixed split-K.""" + """ + 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 + """ + cuda_graph_params: CudaGraphLoraParams = lora_params.get( 'cuda_graph_params') # Get layer-specific parameters @@ -831,12 +843,9 @@ def get_valid_tactics( profile: OptimizationProfile, **kwargs, ) -> List[int]: + # input args are not needed to check valid tactics del inputs, profile, kwargs - k_tiles = max(1, self.input_hidden_size // 64) - return [ - split_k for split_k in _LORA_SPLIT_K_CANDIDATES - if split_k <= k_tiles - ] + return list(_LORA_SPLIT_K_CANDIDATES) def prepare_inputs( self, @@ -844,7 +853,21 @@ def prepare_inputs( lora_params: Dict, cuda_graph_params: CudaGraphLoraParams, ) -> List[torch.Tensor]: - """Copy the LoRA parameters and pack the autotuner tensor inputs.""" + """ + Copy the LoRA parameters and pack the auto-tuner tensor inputs. + + Args: + x: Input tensor + lora_params: LoRA parameters for eager mode + cuda_graph_params: CUDA graph params (also in lora_params) + + Returns: + List of tensor input arguments for runner + + Note: this method is needed because the auto-tuning runner + expects a list of tensors as its input, but the LoRA layer + stores some of them in lora_params and cuda_graph_params. + """ self.lora_params = copy(lora_params) self.cuda_graph_params = copy(cuda_graph_params) @@ -862,7 +885,19 @@ def _prepare_synthetic_inputs( self, inputs: List[torch.Tensor], ) -> List[torch.Tensor]: - """Build one active-adapter problem for the requested token bucket.""" + """ + 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.cuda_graph_params is not None assert self.layer_params is not None @@ -909,6 +944,16 @@ def forward( 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 assert self.cuda_graph_params is not None @@ -920,6 +965,9 @@ def forward( cuda_graph_params.layer_params = {self.layer_key: layer_params} lora_params['cuda_graph_params'] = cuda_graph_params + # The list of tensor input arguments is re-packed + # in the local copies of lora_params and cuda_graph_params + # such that we can invoke _forward_cuda_graph_mode_impl(). x = inputs[0] if len(inputs) > 1: cuda_graph_params.slot_ranks_host = ( From 330afef515016de292c11692495f07eff5ef3c32 Mon Sep 17 00:00:00 2001 From: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> Date: Fri, 31 Jul 2026 02:36:58 -0700 Subject: [PATCH 06/10] Simplify handling of lora_params in auto-tuner Signed-off-by: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> --- tensorrt_llm/_torch/peft/lora/layer.py | 81 +++++++------------ .../_torch/peft/test_lora_autotuner.py | 11 ++- 2 files changed, 38 insertions(+), 54 deletions(-) diff --git a/tensorrt_llm/_torch/peft/lora/layer.py b/tensorrt_llm/_torch/peft/lora/layer.py index 396dac57eaa5..cead3a8d54a3 100644 --- a/tensorrt_llm/_torch/peft/lora/layer.py +++ b/tensorrt_llm/_torch/peft/lora/layer.py @@ -27,7 +27,7 @@ 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, LoraLayerParams +from .cuda_graph_lora_params import CudaGraphLoraParams _FP8_LORA_TMA_ALIGNMENT = 16 @@ -705,7 +705,8 @@ def _forward_cuda_graph_mode( ) runner = self._split_k_runner - runner_inputs = runner.prepare_inputs(x, lora_params, cuda_graph_params) + 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], @@ -811,8 +812,6 @@ def __init__( module_ids=tuple(layer.lora_module_types), ) self.lora_params: Optional[Dict] = None - self.cuda_graph_params: Optional[CudaGraphLoraParams] = None - self.layer_params: Optional[LoraLayerParams] = None self.tuning_config = TuningConfig( dynamic_tensor_specs=(DynamicTensorSpec( 0, @@ -847,39 +846,23 @@ def get_valid_tactics( del inputs, profile, kwargs return list(_LORA_SPLIT_K_CANDIDATES) - def prepare_inputs( - self, - x: torch.Tensor, - lora_params: Dict, - cuda_graph_params: CudaGraphLoraParams, - ) -> List[torch.Tensor]: + def copy_lora_params(self, lora_params: Dict) -> Dict: """ - Copy the LoRA parameters and pack the auto-tuner tensor inputs. + Copy the LoRA parameter hierarchy for this layer. Args: - x: Input tensor - lora_params: LoRA parameters for eager mode - cuda_graph_params: CUDA graph params (also in lora_params) + lora_params: dict to be copied Returns: - List of tensor input arguments for runner - - Note: this method is needed because the auto-tuning runner - expects a list of tensors as its input, but the LoRA layer - stores some of them in lora_params and cuda_graph_params. + Copied lora_params instance """ - self.lora_params = copy(lora_params) - self.cuda_graph_params = copy(cuda_graph_params) - + 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 - self.layer_params = copy(layer_params) - self.cuda_graph_params.layer_params = { - self.layer_key: self.layer_params - } - self.lora_params['cuda_graph_params'] = self.cuda_graph_params - - return [x] + 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, @@ -898,14 +881,16 @@ def _prepare_synthetic_inputs( create the list of tensor input arguments to be used by the auto-tuner's forward. """ - assert self.cuda_graph_params is not None - assert self.layer_params is not None + 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(self.layer_params.d_b_ptrs) - b_prime_ptrs = torch.zeros_like(self.layer_params.d_b_prime_ptrs) + 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): @@ -916,12 +901,12 @@ def _prepare_synthetic_inputs( b_prime_ptrs[module_idx, 0] = lora_b.data_ptr() keepalive.extend((lora_a, lora_b)) - slot_counts = torch.zeros_like(self.cuda_graph_params.slot_counts) + slot_counts = torch.zeros_like(cuda_graph_params.slot_counts) slot_counts[0] = num_tokens - slot_ranks = torch.zeros_like(self.cuda_graph_params.slot_ranks) + slot_ranks = torch.zeros_like(cuda_graph_params.slot_ranks) slot_ranks[0] = self.max_rank slot_offsets_full = torch.zeros_like( - self.cuda_graph_params.slot_offsets_full) + cuda_graph_params.slot_offsets_full) slot_offsets_full[1:] = num_tokens return [ @@ -932,8 +917,8 @@ def _prepare_synthetic_inputs( b_ptrs, b_prime_ptrs, torch.arange(num_tokens, device=token_carrier.device), - self.layer_params.d_output_sizes, - self.layer_params.d_output_sizes_offset, + layer_params.d_output_sizes, + layer_params.d_output_sizes_offset, ] + keepalive def forward( @@ -956,18 +941,14 @@ def forward( """ del kwargs assert self.lora_params is not None - assert self.cuda_graph_params is not None - assert self.layer_params is not None - - lora_params = copy(self.lora_params) - cuda_graph_params = copy(self.cuda_graph_params) - layer_params = copy(self.layer_params) - cuda_graph_params.layer_params = {self.layer_key: layer_params} - lora_params['cuda_graph_params'] = cuda_graph_params - - # The list of tensor input arguments is re-packed - # in the local copies of lora_params and cuda_graph_params - # such that we can invoke _forward_cuda_graph_mode_impl(). + lora_params = self.copy_lora_params(self.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 + + # The list of tensor input arguments is re-packed into + # the local lora_params copy such that we can invoke + # LoraLayer's _forward_cuda_graph_mode_impl(). x = inputs[0] if len(inputs) > 1: cuda_graph_params.slot_ranks_host = ( diff --git a/tests/unittest/_torch/peft/test_lora_autotuner.py b/tests/unittest/_torch/peft/test_lora_autotuner.py index 5be2f6a0d8e9..9227436e90ee 100644 --- a/tests/unittest/_torch/peft/test_lora_autotuner.py +++ b/tests/unittest/_torch/peft/test_lora_autotuner.py @@ -56,10 +56,10 @@ def test_lora_split_k_runner_uses_token_buckets(): assert profile[0][0] == 4 -def test_lora_split_k_runner_prunes_splits_larger_than_k_tiles(): +def test_lora_split_k_runner_returns_all_candidates(): runner = _make_runner(input_hidden_size=256) - assert runner.get_valid_tactics([], None) == [1, 2, 4] + assert runner.get_valid_tactics([], None) == [1, 2, 4, 8, 16] def test_lora_layer_reuses_runner_across_cuda_graph_warmups(monkeypatch): @@ -168,17 +168,20 @@ def test_lora_autotuner_hook_builds_single_active_slot(): dtype=runner.dtype, device="cuda", ) - runner.layer_params = SimpleNamespace( + 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"), ) - runner.cuda_graph_params = SimpleNamespace( + 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] From b9ca32a9d63adabd88e5e41952dbcfb8c84c83d1 Mon Sep 17 00:00:00 2001 From: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> Date: Fri, 31 Jul 2026 02:51:15 -0700 Subject: [PATCH 07/10] Drop trivial auto-tuner tests Signed-off-by: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> --- tests/unittest/_torch/peft/test_lora_autotuner.py | 13 ------------- 1 file changed, 13 deletions(-) diff --git a/tests/unittest/_torch/peft/test_lora_autotuner.py b/tests/unittest/_torch/peft/test_lora_autotuner.py index 9227436e90ee..ab919ab36c5f 100644 --- a/tests/unittest/_torch/peft/test_lora_autotuner.py +++ b/tests/unittest/_torch/peft/test_lora_autotuner.py @@ -32,13 +32,6 @@ def _make_runner( ) -def test_lora_split_k_runner_identity_is_layer_specific(): - runner = _make_runner(layer_idx=3) - other_layer_runner = _make_runner(layer_idx=4) - - assert runner.unique_id() != other_layer_runner.unique_id() - - def test_lora_split_k_runner_uses_token_buckets(): runner = _make_runner() spec = runner.tuning_config.dynamic_tensor_specs[0] @@ -56,12 +49,6 @@ def test_lora_split_k_runner_uses_token_buckets(): assert profile[0][0] == 4 -def test_lora_split_k_runner_returns_all_candidates(): - runner = _make_runner(input_hidden_size=256) - - assert runner.get_valid_tactics([], None) == [1, 2, 4, 8, 16] - - def test_lora_layer_reuses_runner_across_cuda_graph_warmups(monkeypatch): layer = lora_layer.LoraLayer( [lora_layer.LoraModuleType.ATTENTION_Q], From 48691764410df82616a291346d025dd446101878 Mon Sep 17 00:00:00 2001 From: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> Date: Wed, 5 Aug 2026 23:48:55 -0700 Subject: [PATCH 08/10] Remove redundant lora_params copy Signed-off-by: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> --- tensorrt_llm/_torch/peft/lora/layer.py | 17 ++++++++--------- 1 file changed, 8 insertions(+), 9 deletions(-) diff --git a/tensorrt_llm/_torch/peft/lora/layer.py b/tensorrt_llm/_torch/peft/lora/layer.py index cead3a8d54a3..06405207b07e 100644 --- a/tensorrt_llm/_torch/peft/lora/layer.py +++ b/tensorrt_llm/_torch/peft/lora/layer.py @@ -697,7 +697,7 @@ def _forward_cuda_graph_mode( self._split_k_runner = _LoraGroupedGemmRunner( layer=self, layer_idx=layer_idx, - input_hidden_size=x.shape[1], + 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), @@ -941,16 +941,15 @@ def forward( """ del kwargs assert self.lora_params is not None - lora_params = self.copy_lora_params(self.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 - - # The list of tensor input arguments is re-packed into - # the local lora_params copy such that we can invoke - # LoraLayer's _forward_cuda_graph_mode_impl(). + 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 cuda_graph_params.slot_ranks_host = ( cuda_graph_params.slot_ranks_host.clone()) cuda_graph_params.slot_ranks_host.zero_() From bff6efa7c89a80420b21533fcdf601b22606910e Mon Sep 17 00:00:00 2001 From: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> Date: Thu, 6 Aug 2026 02:28:59 -0700 Subject: [PATCH 09/10] Test for default autotuner tactic, check that split-K value is passed to grouped GEMM properly Signed-off-by: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> --- .../_torch/peft/test_lora_autotuner.py | 68 +++++++++++++------ 1 file changed, 48 insertions(+), 20 deletions(-) diff --git a/tests/unittest/_torch/peft/test_lora_autotuner.py b/tests/unittest/_torch/peft/test_lora_autotuner.py index ab919ab36c5f..564a33a3e41a 100644 --- a/tests/unittest/_torch/peft/test_lora_autotuner.py +++ b/tests/unittest/_torch/peft/test_lora_autotuner.py @@ -49,6 +49,23 @@ def test_lora_split_k_runner_uses_token_buckets(): 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_lora_layer_reuses_runner_across_cuda_graph_warmups(monkeypatch): layer = lora_layer.LoraLayer( [lora_layer.LoraModuleType.ATTENTION_Q], @@ -61,19 +78,19 @@ def test_lora_layer_reuses_runner_across_cuda_graph_warmups(monkeypatch): ) class FakeLayerParams: - def __init__(self): - self.d_b_ptrs = torch.tensor([1]) - self.d_b_prime_ptrs = torch.tensor([2]) + 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.layer_info = {layer_key: object()} - self.layer_params = {layer_key: FakeLayerParams()} 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]) @@ -99,31 +116,43 @@ def choose_one(self, custom_op, runners, tuning_config, inputs): monkeypatch.setattr(lora_layer.AutoTuner, "get", staticmethod(lambda: FakeTuner())) - forwarded_params = [] + parameter_fill_calls = [] - def fake_forward_impl(x, runner_lora_params, forwarded_layer_idx, split_k): - forwarded_params.append((runner_lora_params, split_k)) - assert forwarded_layer_idx == layer_idx - return x + 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(layer, "_forward_cuda_graph_mode_impl", fake_forward_impl) + 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) - assert layer._forward_cuda_graph_mode(x, lora_params, layer_idx) is x + 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 len(forwarded_params) == 2 - assert forwarded_params[0][0] is not forwarded_params[1][0] + assert operator_split_ks == [1, 1] + assert warmup_lora_params[0] is not warmup_lora_params[1] - runner_lora_params = forwarded_params[-1][0] + 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 @@ -133,12 +162,11 @@ def fake_forward_impl(x, runner_lora_params, forwarded_layer_idx, split_k): synthetic_inputs = runner._prepare_synthetic_inputs([torch.empty(2, 256)]) runner(synthetic_inputs, tactic=2) + assert operator_split_ks == [1, 1, 2] - synthetic_lora_params = forwarded_params[-1][0] - synthetic_cuda_graph_params = synthetic_lora_params["cuda_graph_params"] - synthetic_layer_params = synthetic_cuda_graph_params.get_layer_params(layer_key) - assert synthetic_cuda_graph_params.slot_counts is synthetic_inputs[1] - assert synthetic_layer_params.d_b_ptrs is synthetic_inputs[4] + 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 From 063e5b325cb1ab491c48c60b0311c72f037a6da8 Mon Sep 17 00:00:00 2001 From: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> Date: Mon, 17 Aug 2026 06:22:32 -0700 Subject: [PATCH 10/10] Integrate FP8 LoRA with split-K autotuning Signed-off-by: Alessio Netti <26897207+AlessioNetti@users.noreply.github.com> --- tensorrt_llm/_torch/peft/lora/layer.py | 9 ++-- .../_torch/peft/test_lora_autotuner.py | 53 ++++++++++++++++++- 2 files changed, 57 insertions(+), 5 deletions(-) diff --git a/tensorrt_llm/_torch/peft/lora/layer.py b/tensorrt_llm/_torch/peft/lora/layer.py index 06405207b07e..4d7e87de1935 100644 --- a/tensorrt_llm/_torch/peft/lora/layer.py +++ b/tensorrt_llm/_torch/peft/lora/layer.py @@ -950,10 +950,11 @@ def forward( 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 - 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 + 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, diff --git a/tests/unittest/_torch/peft/test_lora_autotuner.py b/tests/unittest/_torch/peft/test_lora_autotuner.py index 564a33a3e41a..bcd661d3c38d 100644 --- a/tests/unittest/_torch/peft/test_lora_autotuner.py +++ b/tests/unittest/_torch/peft/test_lora_autotuner.py @@ -13,6 +13,7 @@ def _make_runner( layer_idx: int = 3, input_hidden_size: int = 256, + dtype: torch.dtype = torch.float16, ) -> lora_layer._LoraGroupedGemmRunner: layer = lora_layer.LoraLayer( [ @@ -28,7 +29,7 @@ def _make_runner( max_rank=16, max_lora_size=4, problem_count=8, - dtype=torch.float16, + dtype=dtype, ) @@ -66,6 +67,56 @@ def fake_forward_impl(x, lora_params, layer_idx, split_k): 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],