[LoRA] Optimize TrtLlmLoRAExperts (#48759)

Signed-off-by: Jee Jee Li <jeejeelee@inferact.ai>
This commit is contained in:
Jee Jee Li
2026-07-18 10:26:14 +08:00
committed by GitHub
parent 425c4eafb0
commit da64db78b9
2 changed files with 221 additions and 41 deletions
@@ -46,6 +46,7 @@ class LoRAExpertsMixin:
num_tokens: int,
top_k_num: int,
add_inputs: bool = True,
swap_w13_slices: bool = False,
) -> tuple[
torch.Tensor | None,
torch.Tensor | None,
@@ -53,6 +54,7 @@ class LoRAExpertsMixin:
torch.Tensor | None,
]:
w13_lora_a_stacked = lora_context.w13_lora_a_stacked
w13_lora_b_stacked = lora_context.w13_lora_b_stacked
if lora_context.enable_moe_shared_loras:
# w13 lora_A is shared across experts (collapsed expert-dim 1);
# broadcast to local_num_experts via a stride-0 view. The kernel
@@ -61,11 +63,19 @@ class LoRAExpertsMixin:
a.expand(-1, lora_context.local_num_experts, -1, -1)
for a in w13_lora_a_stacked
)
if swap_w13_slices:
# The expand kernel writes slice j into the j-th half of y's last
# dim. Reversing the (gate, up) slice tuples makes it emit
# [up, gate] order directly -- used by the FlashInfer trtllm path,
# whose SwiGLU expects the up half first, to avoid an out-of-place
# concat swap afterwards.
w13_lora_a_stacked = w13_lora_a_stacked[::-1]
w13_lora_b_stacked = w13_lora_b_stacked[::-1]
return lora_context.punica_wrapper.add_lora_w13(
y,
x,
w13_lora_a_stacked,
lora_context.w13_lora_b_stacked,
w13_lora_b_stacked,
topk_ids,
topk_weights,
expert_map,
@@ -31,6 +31,7 @@ from vllm.model_executor.layers.fused_moe.config import (
FusedMoEQuantConfig,
RoutingMethodType,
)
from vllm.model_executor.layers.fused_moe.experts.lora_context import MoELoRAContext
from vllm.model_executor.layers.fused_moe.experts.lora_experts_mixin import (
LoRAExpertsMixin,
)
@@ -41,9 +42,75 @@ from vllm.model_executor.layers.fused_moe.utils import (
trtllm_moe_pack_topk_ids_weights,
)
from vllm.platforms import current_platform
from vllm.triton_utils import tl, triton
from vllm.utils.flashinfer import has_flashinfer_trtllm_fused_moe
@triton.jit
def _unpermute_activation_kernel(
act_ptr, # act_permuted: (num_permuted, num_cols)
idx_ptr, # idx_map: (num_rows,), values in [0, num_permuted) or -1
out_ptr, # out: (num_rows, num_cols)
num_cols,
stride_ar,
stride_or,
BLOCK_I: tl.constexpr,
):
row = tl.program_id(0)
col_offs = tl.program_id(1) * BLOCK_I + tl.arange(0, BLOCK_I)
col_mask = col_offs < num_cols
idx = tl.load(idx_ptr + row)
out_ptrs = out_ptr + row * stride_or + col_offs
if idx >= 0:
vals = tl.load(act_ptr + idx * stride_ar + col_offs, mask=col_mask, other=0.0)
tl.store(out_ptrs, vals, mask=col_mask)
else:
zeros = tl.zeros((BLOCK_I,), dtype=out_ptr.dtype.element_ty)
tl.store(out_ptrs, zeros, mask=col_mask)
@triton.jit
def _finalize_lora_kernel(
gemm2_ptr, # (num_permuted, K) base FC2 output, permuted, unweighted
weight_ptr, # (num_tokens * top_k,) routing weights (expanded order)
idx_ptr, # (num_tokens * top_k,) expanded_idx -> permuted_idx or -1
delta_ptr, # (num_tokens, top_k, K) W2 LoRA delta, already routing-weighted
out_ptr, # (num_tokens, K)
K,
stride_g0,
stride_d0,
stride_d1,
stride_o0,
scale,
TOP_K: tl.constexpr,
BLOCK_K: tl.constexpr,
):
token = tl.program_id(0)
col = tl.program_id(1) * BLOCK_K + tl.arange(0, BLOCK_K)
mask = col < K
acc_base = tl.zeros((BLOCK_K,), dtype=tl.float32)
acc_delta = tl.zeros((BLOCK_K,), dtype=tl.float32)
for k in tl.static_range(TOP_K):
eid = token * TOP_K + k
pidx = tl.load(idx_ptr + eid)
if pidx >= 0:
w = tl.load(weight_ptr + eid).to(tl.float32)
base = tl.load(gemm2_ptr + pidx * stride_g0 + col, mask=mask, other=0.0).to(
tl.float32
)
acc_base += w * base
acc_delta += tl.load(
delta_ptr + token * stride_d0 + k * stride_d1 + col, mask=mask, other=0.0
).to(tl.float32)
out = acc_base * scale + acc_delta
tl.store(
out_ptr + token * stride_o0 + col, out.to(out_ptr.dtype.element_ty), mask=mask
)
class _TrtLlmLoRAExpertsBase(LoRAExpertsMixin, mk.FusedMoEExpertsModular):
"""LoRA-aware trtllm MoE experts"""
@@ -107,7 +174,8 @@ class _TrtLlmLoRAExpertsBase(LoRAExpertsMixin, mk.FusedMoEExpertsModular):
return E, M, N, K, topk
def finalize_weight_and_reduce_impl(self) -> mk.TopKWeightAndReduce:
# do_finalize=True: the kernel already does moe_sum, so this is a No-Op
# apply() writes the fully finalized result into `output` (fused base
# finalize + W2 LoRA reduction), so this is a No-Op.
return TopKWeightAndReduceNoOP()
@staticmethod
@@ -146,9 +214,13 @@ class _TrtLlmLoRAExpertsBase(LoRAExpertsMixin, mk.FusedMoEExpertsModular):
) -> list[torch.Tensor]:
"""Call the dtype-specific trtllm_*_routed_moe and return list[Tensor].
Return contract (do_finalize=True):
gemm1_lora_delta is None -> [output]
otherwise -> [output, expanded_idx_to_permuted_idx,
The LoRA path always sets gemm1_lora_delta and runs with
do_finalize=False so the base finalize can be fused with the W2 LoRA
reduction (see _finalize_with_w2_lora). Return contract:
gemm1_lora_delta is None -> [output] (do_finalize=True)
otherwise -> [gemm2_output(permuted, unweighted),
expert_weights,
expanded_idx_to_permuted_idx,
gemm1_activation_output(permuted)]
"""
raise NotImplementedError
@@ -178,6 +250,27 @@ class _TrtLlmLoRAExpertsBase(LoRAExpertsMixin, mk.FusedMoEExpertsModular):
intermediate_size = self.intermediate_size_per_partition
K = output.size(1)
# Routing is computed outside the MoE; pack it into the
# (eid<<16)|w.bf16 format the routed API expects.
packed_topk_ids = trtllm_moe_pack_topk_ids_weights(topk_ids, topk_weights)
# ---- Base-model fast path ----
# When no token in the batch selects a LoRA adapter, skip the LoRA machinery
# and run the plain base MoE with do_finalize=True, which writes the finalized
# result straight into `output`.
if self._batch_has_no_lora(lora_context):
self.invoke_routed_moe(
hidden_states=hidden_states,
w1=w1,
w2=w2,
packed_topk_ids=packed_topk_ids,
gemm1_lora_delta=None, # without LoRA, no delta
global_num_experts=global_num_experts,
a1q_scale=a1q_scale,
output=output,
)
return
# The LoRA tile-config heuristic (try_get_optimal_moe_config) unpacks
# w1/w2 as standard 3D MoE weights, but flashinfer stores shuffled
# 4D BlockMajorK weights. add_lora_w13/add_lora_w2 only read .shape
@@ -195,14 +288,13 @@ class _TrtLlmLoRAExpertsBase(LoRAExpertsMixin, mk.FusedMoEExpertsModular):
dtype=torch.bfloat16,
)
# Routing is computed outside the MoE; pack it into the
# (eid<<16)|w.bf16 format the routed API expects.
packed_topk_ids = trtllm_moe_pack_topk_ids_weights(topk_ids, topk_weights)
# ---- 1) W13 LoRA delta -> gemm1_lora_delta (bf16, [T, top_k, 2I]) ----
gemm1_lora_delta = None
w13_meta = (None, None, None, None)
# zeros (not empty): under EP the punica expand kernel only writes
# slots whose expert is local to this rank; non-local (token, top_k)
# slots must stay 0 so they contribute no bias when fed to flashinfer.
gemm1_lora_delta = torch.zeros(
num_tokens,
top_k,
@@ -221,6 +313,10 @@ class _TrtLlmLoRAExpertsBase(LoRAExpertsMixin, mk.FusedMoEExpertsModular):
# add_inputs=False: write the pure delta only (the base is fused in
# by the kernel) and do NOT multiply by the routing weight (it is a
# pre-SwiGLU bias).
# swap_w13_slices=True: apply_w13_lora writes in vLLM's w13 order
# (gate=w1 first, up=w3 second), but FlashInfer's gemm1_lora_delta
# expects [up, gate]; reversing the slices emits that order directly,
# avoiding an out-of-place concat swap.
w13_meta = self.apply_w13_lora(
lora_context,
y=gemm1_lora_delta,
@@ -233,18 +329,7 @@ class _TrtLlmLoRAExpertsBase(LoRAExpertsMixin, mk.FusedMoEExpertsModular):
num_tokens=num_tokens,
top_k_num=top_k,
add_inputs=False,
)
# apply_w13_lora writes the delta in vLLM's w13 order (gate=w1 first,
# up=w3 second), but FlashInfer's gemm1_lora_delta expects the halves
# in [up, gate] order. Swap them so the delta lands on the matching
# SwiGLU branch.
gemm1_lora_delta = torch.cat(
[
gemm1_lora_delta[..., intermediate_size:],
gemm1_lora_delta[..., :intermediate_size],
],
dim=-1,
swap_w13_slices=True,
)
# ---- 2) Call the routed flashinfer kernel ----
@@ -259,8 +344,14 @@ class _TrtLlmLoRAExpertsBase(LoRAExpertsMixin, mk.FusedMoEExpertsModular):
output=output,
)
# ---- 3) W2 LoRA (computed out of kernel) ----
expanded_idx_to_permuted_idx = ret[1]
gemm1_act_permuted = ret[2] # [max_padded, I], post-act
# do_finalize=False: flashinfer returns the *unfinalized* base output.
# ret = [gemm2_output(permuted, unweighted),
# expert_weights, expanded_idx_to_permuted_idx,
# gemm1_activation_output(permuted)]
gemm2_permuted = ret[0]
expert_weights = ret[1]
expanded_idx_to_permuted_idx = ret[2]
gemm1_act_permuted = ret[3] # [max_padded, I], post-act
act = self._unpermute_activation(
gemm1_act_permuted,
expanded_idx_to_permuted_idx,
@@ -298,11 +389,37 @@ class _TrtLlmLoRAExpertsBase(LoRAExpertsMixin, mk.FusedMoEExpertsModular):
top_k_num=top_k,
add_inputs=False,
)
# The base output is already finalized (routing-weighted + summed over
# top_k); the W2 delta is likewise already routing-weighted, so sum it
# over top_k and add.
# TODO(verify): if routed_scaling_factor is not None, scale to match.
output.add_(w2_delta.sum(dim=1))
# ---- 4) Fused finalize: reduce the base path over top_k (with routing
# weights) and add the already-weighted W2 delta, in a single kernel.
# This replaces flashinfer's internal finalize launch plus a separate
# w2_delta.sum(dim=1) + add_.
self._finalize_with_w2_lora(
output,
gemm2_permuted,
expert_weights,
expanded_idx_to_permuted_idx,
w2_delta,
num_tokens,
top_k,
scale=1.0,
)
@staticmethod
def _batch_has_no_lora(lora_context: MoELoRAContext) -> bool:
"""True when no token in the batch selects a LoRA adapter.
Mirrors the no-lora fast path in
``PunicaWrapperGPU.add_lora_fused_moe``: the punica kernel metadata
carries a CPU ``no_lora_flag`` computed once per forward from the
token->LoRA mapping. Reading it is a host-only check (no device sync),
and under CUDA graphs the branch is frozen at capture time against the
graph's ``has_lora`` dispatch key, so it stays correct on replay.
"""
meta = getattr(lora_context.punica_wrapper, "token_mapping_meta", None)
if meta is None:
return False
flag = meta.no_lora_flag_cpu
return bool(flag.numel() == 1 and flag.item())
@staticmethod
def _unpermute_activation(
@@ -315,13 +432,63 @@ class _TrtLlmLoRAExpertsBase(LoRAExpertsMixin, mk.FusedMoEExpertsModular):
"""Permuted FC1 activation -> (num_tokens*top_k, I).
expanded_idx = token*top_k + k; idx_map[expanded_idx] = permuted_idx or -1.
TODO optimize these operations
Fused gather + drop-masking: each output row copies the matching
permuted row, or is zeroed when idx_map < 0.
"""
num_rows = num_tokens * top_k
out = torch.empty(
(num_rows, intermediate_size),
dtype=act_permuted.dtype,
device=act_permuted.device,
)
BLOCK_I = 1024
grid = (num_rows, triton.cdiv(intermediate_size, BLOCK_I))
_unpermute_activation_kernel[grid](
act_permuted,
idx_map,
out,
intermediate_size,
act_permuted.stride(0),
out.stride(0),
BLOCK_I=BLOCK_I,
)
return out
valid = idx_map >= 0
safe_idx = idx_map.clamp_min(0).long()
gathered = act_permuted[safe_idx]
return gathered * valid.unsqueeze(1).to(act_permuted.dtype)
@staticmethod
def _finalize_with_w2_lora(
output: torch.Tensor,
gemm2_permuted: torch.Tensor,
expert_weights: torch.Tensor,
idx_map: torch.Tensor,
w2_delta: torch.Tensor,
num_tokens: int,
top_k: int,
scale: float = 1.0,
) -> None:
"""Fused base finalize + W2 LoRA reduction, written into ``output``.
For each token: sum the routing-weighted permuted base rows over top_k
(``expert_weights`` in expanded order, ``idx_map < 0`` dropped), scale by
``scale``, and add the already-weighted ``w2_delta`` reduced over top_k.
"""
K = gemm2_permuted.size(1)
BLOCK_K = 512
grid = (num_tokens, triton.cdiv(K, BLOCK_K))
_finalize_lora_kernel[grid](
gemm2_permuted,
expert_weights.reshape(-1),
idx_map,
w2_delta,
output,
K,
gemm2_permuted.stride(0),
w2_delta.stride(0),
w2_delta.stride(1),
output.stride(0),
scale,
TOP_K=top_k,
BLOCK_K=BLOCK_K,
)
# BF16 unquantized trtllm MoE + LoRA
@@ -359,10 +526,12 @@ class TrtLlmBf16LoRAExperts(_TrtLlmLoRAExpertsBase):
) -> list[torch.Tensor]:
import flashinfer
# Unlike the fp8/mxint4 routed APIs, trtllm_bf16_routed_moe has no
# `output=` kwarg: it returns the finalized tensor (or a list whose
# [0] is it when gemm1_lora_delta is set). Copy it into the caller's
# buffer so the modular-kernel output plumbing sees the result.
# With gemm1_lora_delta set (the LoRA path) run do_finalize=False and
# return the unfinalized permuted base output so apply() can fuse the
# finalize with the W2 LoRA reduction (see _finalize_with_w2_lora).
# Without a delta (base path), run do_finalize=True and hand flashinfer
# the caller's buffer via output= so it finalizes in place -- no copy.
do_finalize = gemm1_lora_delta is None
ret = flashinfer.fused_moe.trtllm_bf16_routed_moe(
topk_ids=packed_topk_ids,
hidden_states=hidden_states,
@@ -378,10 +547,11 @@ class TrtLlmBf16LoRAExperts(_TrtLlmLoRAExpertsBase):
local_num_experts=self.local_num_experts,
routed_scaling_factor=None,
routing_method_type=self.routing_method_type,
do_finalize=True,
do_finalize=do_finalize,
output=output if do_finalize else None,
)
if isinstance(ret, (list, tuple)):
output.copy_(ret[0])
if not do_finalize:
# [gemm2_output, expert_weights, expanded_idx, gemm1_activation]
return list(ret)
output.copy_(ret)
# do_finalize=True finalized directly into `output`.
return [output]