(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_cuda = is_cuda()
|
||||||
_is_hip = is_hip()
|
_is_hip = is_hip()
|
||||||
_is_hip = is_hip()
|
|
||||||
_is_xpu = is_xpu()
|
_is_xpu = is_xpu()
|
||||||
|
|
||||||
if _is_cuda or _is_hip or _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
|
_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(
|
def _naive_moe_lora_align_block_size(
|
||||||
topk_ids: torch.Tensor,
|
topk_ids: torch.Tensor,
|
||||||
seg_indptr: torch.Tensor,
|
seg_indptr: torch.Tensor,
|
||||||
@@ -288,7 +181,6 @@ class LoRAInfo:
|
|||||||
num_experts: int
|
num_experts: int
|
||||||
experts_shared_outer_loras: bool = False
|
experts_shared_outer_loras: bool = False
|
||||||
cg_buffers: dict | None = None
|
cg_buffers: dict | None = None
|
||||||
cg_buffers: dict | None = None
|
|
||||||
|
|
||||||
fully_sharded: bool = False
|
fully_sharded: bool = False
|
||||||
tp_size: int = 1
|
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
|
from sglang.srt.model_executor.cuda_graph_runner import get_capture_lora_variant
|
||||||
|
|
||||||
# Record LoRA kernels for lora graph; skip for nolora graph.
|
# Record LoRA kernels for lora graph; skip for nolora graph.
|
||||||
has_active_lora = get_capture_lora_variant() != "nolora"
|
if get_capture_lora_variant() == "nolora":
|
||||||
else:
|
return
|
||||||
num_loras = len(lora_info.lora_ranks)
|
|
||||||
has_active_lora = (
|
if lora_info is None or lora_info.max_lora_rank == 0:
|
||||||
(
|
|
||||||
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:
|
|
||||||
return
|
return
|
||||||
|
|
||||||
M, top_k, gate_up_dim = intermediate_cache.shape
|
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:
|
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)
|
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:
|
if lora_info.fully_sharded and lora_info.tp_size > 1:
|
||||||
shard_size = lora_info.hidden_size // lora_info.tp_size
|
shard_size = lora_info.hidden_size // lora_info.tp_size
|
||||||
offset = shard_size * lora_info.tp_rank
|
offset = shard_size * lora_info.tp_rank
|
||||||
|
|||||||
Reference in New Issue
Block a user