Skip to content
Merged
263 changes: 260 additions & 3 deletions tensorrt_llm/_torch/peft/lora/layer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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:
Expand All @@ -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,
Expand Down Expand Up @@ -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.

Expand Down
23 changes: 22 additions & 1 deletion tensorrt_llm/_torch/pyexecutor/model_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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 = {}
Expand Down
Loading
Loading