(2/n - prefill optimize)perf(lora): remove GPU-CPU sync barrier (.item()) in MoE LoRA path and remove duplicate code (#24246)
Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
co-authored by
Cursor
parent
2f7d99b7f7
commit
2b769d37a4
@@ -34,7 +34,6 @@ from sglang.srt.utils import is_cuda, is_hip, is_xpu, next_power_of_2
|
||||
|
||||
_is_cuda = is_cuda()
|
||||
_is_hip = is_hip()
|
||||
_is_hip = is_hip()
|
||||
_is_xpu = is_xpu()
|
||||
|
||||
if _is_cuda or _is_hip or _is_xpu:
|
||||
@@ -64,112 +63,6 @@ def _get_moe_lora_block_config(max_lora_rank: int) -> dict:
|
||||
_SPARSITY_FACTOR = 8
|
||||
|
||||
|
||||
def _naive_moe_lora_align_block_size(
|
||||
topk_ids: torch.Tensor,
|
||||
seg_indptr: torch.Tensor,
|
||||
req_to_lora: torch.Tensor,
|
||||
num_experts: int,
|
||||
block_size_m: int,
|
||||
max_loras: int,
|
||||
max_num_tokens_padded: int,
|
||||
max_num_m_blocks: int,
|
||||
adapter_enabled: torch.Tensor,
|
||||
device: torch.device,
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
"""Construct LoRA token-expert alignment on CPU for small batches.
|
||||
|
||||
When the number of tokens is very small, the overhead of launching the
|
||||
CUDA-based moe_lora_align_block_size kernel exceeds the actual
|
||||
computation. This function builds the same data structures using simple
|
||||
Python loops on CPU and transfers the result to GPU in one shot.
|
||||
"""
|
||||
M, top_k = topk_ids.shape
|
||||
num_valid_tokens = M * top_k
|
||||
|
||||
sorted_token_ids = torch.full(
|
||||
(max_loras * max_num_tokens_padded,),
|
||||
num_valid_tokens,
|
||||
dtype=torch.int32,
|
||||
)
|
||||
expert_ids_out = torch.full((max_loras * max_num_m_blocks,), -1, dtype=torch.int32)
|
||||
num_tokens_post_padded = torch.zeros(max_loras, dtype=torch.int32)
|
||||
|
||||
seg_indptr_list = seg_indptr.cpu().tolist()
|
||||
req_to_lora_list = req_to_lora.cpu().tolist()
|
||||
topk_ids_list = topk_ids.cpu().tolist()
|
||||
adapter_enabled_list = adapter_enabled.cpu().tolist()
|
||||
|
||||
for lora_id in range(max_loras):
|
||||
if not adapter_enabled_list[lora_id]:
|
||||
continue
|
||||
|
||||
pairs: list[tuple[int, int]] = []
|
||||
for seg_idx in range(len(seg_indptr_list) - 1):
|
||||
if req_to_lora_list[seg_idx] != lora_id:
|
||||
continue
|
||||
start = seg_indptr_list[seg_idx]
|
||||
end = seg_indptr_list[seg_idx + 1]
|
||||
for m in range(start, end):
|
||||
for k in range(top_k):
|
||||
pairs.append((topk_ids_list[m][k], m * top_k + k))
|
||||
|
||||
if not pairs:
|
||||
continue
|
||||
|
||||
pairs.sort()
|
||||
|
||||
base_t = lora_id * max_num_tokens_padded
|
||||
base_e = lora_id * max_num_m_blocks
|
||||
pos = 0
|
||||
block_idx = 0
|
||||
i = 0
|
||||
while i < len(pairs):
|
||||
cur_expert = pairs[i][0]
|
||||
group_start = pos
|
||||
while i < len(pairs) and pairs[i][0] == cur_expert:
|
||||
sorted_token_ids[base_t + pos] = pairs[i][1]
|
||||
pos += 1
|
||||
i += 1
|
||||
group_len = pos - group_start
|
||||
padded_len = ((group_len + block_size_m - 1) // block_size_m) * block_size_m
|
||||
num_blocks = padded_len // block_size_m
|
||||
for b in range(num_blocks):
|
||||
expert_ids_out[base_e + block_idx + b] = cur_expert
|
||||
block_idx += num_blocks
|
||||
pos = group_start + padded_len
|
||||
|
||||
num_tokens_post_padded[lora_id] = pos
|
||||
|
||||
return (
|
||||
sorted_token_ids.to(device),
|
||||
expert_ids_out.to(device),
|
||||
num_tokens_post_padded.to(device),
|
||||
)
|
||||
|
||||
|
||||
def _get_moe_lora_block_config(max_lora_rank: int) -> dict:
|
||||
"""Compute rank-aware block sizes for MoE LoRA kernels.
|
||||
|
||||
Shrink: output dim is the rank -> cap BLOCK_SIZE_N to avoid waste.
|
||||
Expand: input dim is the rank -> cap BLOCK_SIZE_K similarly.
|
||||
"""
|
||||
if max_lora_rank <= 0:
|
||||
rank_pow2 = 64
|
||||
else:
|
||||
rank_pow2 = next_power_of_2(max_lora_rank)
|
||||
|
||||
shrink_n = min(64, rank_pow2)
|
||||
expand_k = max(16, min(64, rank_pow2))
|
||||
|
||||
return {
|
||||
"shrink_block_size_n": shrink_n,
|
||||
"expand_block_size_k": expand_k,
|
||||
}
|
||||
|
||||
|
||||
_SPARSITY_FACTOR = 8
|
||||
|
||||
|
||||
def _naive_moe_lora_align_block_size(
|
||||
topk_ids: torch.Tensor,
|
||||
seg_indptr: torch.Tensor,
|
||||
@@ -288,7 +181,6 @@ class LoRAInfo:
|
||||
num_experts: int
|
||||
experts_shared_outer_loras: bool = False
|
||||
cg_buffers: dict | None = None
|
||||
cg_buffers: dict | None = None
|
||||
|
||||
fully_sharded: bool = False
|
||||
tp_size: int = 1
|
||||
@@ -447,18 +339,10 @@ def _add_lora_gate_up_delta(
|
||||
from sglang.srt.model_executor.cuda_graph_runner import get_capture_lora_variant
|
||||
|
||||
# Record LoRA kernels for lora graph; skip for nolora graph.
|
||||
has_active_lora = get_capture_lora_variant() != "nolora"
|
||||
else:
|
||||
num_loras = len(lora_info.lora_ranks)
|
||||
has_active_lora = (
|
||||
(
|
||||
lora_info.adapter_enabled[:num_loras]
|
||||
* (lora_info.lora_ranks > 0).to(lora_info.adapter_enabled.dtype)
|
||||
)
|
||||
.any()
|
||||
.item()
|
||||
)
|
||||
if not has_active_lora or lora_info is None or lora_info.max_lora_rank == 0:
|
||||
if get_capture_lora_variant() == "nolora":
|
||||
return
|
||||
|
||||
if lora_info is None or lora_info.max_lora_rank == 0:
|
||||
return
|
||||
|
||||
M, top_k, gate_up_dim = intermediate_cache.shape
|
||||
@@ -565,11 +449,6 @@ def _add_lora_down_delta(
|
||||
if lora_info.experts_shared_outer_loras and not lora_info.lora_use_virtual_experts:
|
||||
down_lora_b = down_lora_b.expand(-1, lora_info.num_experts, -1, -1)
|
||||
|
||||
if lora_info.fully_sharded and lora_info.tp_size > 1:
|
||||
shard_size = lora_info.hidden_size // lora_info.tp_size
|
||||
offset = shard_size * lora_info.tp_rank
|
||||
else:
|
||||
offset = 0
|
||||
if lora_info.fully_sharded and lora_info.tp_size > 1:
|
||||
shard_size = lora_info.hidden_size // lora_info.tp_size
|
||||
offset = shard_size * lora_info.tp_rank
|
||||
|
||||
Reference in New Issue
Block a user