feat(lora): enable csgmv backend with virtual experts for MoE LoRA (#24007)
This commit is contained in:
@@ -203,6 +203,18 @@ class ChunkedSgmvLoRABackend(BaseLoRABackend):
|
|||||||
chunk_size = 16
|
chunk_size = 16
|
||||||
return min(self.max_chunk_size, chunk_size)
|
return min(self.max_chunk_size, chunk_size)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _build_req_seg_indptr(forward_batch: ForwardBatch) -> torch.Tensor:
|
||||||
|
"""Build per-request cumulative token boundaries on CPU (pinned)."""
|
||||||
|
bs = forward_batch.batch_size
|
||||||
|
if forward_batch.forward_mode.is_decode():
|
||||||
|
indptr = torch.arange(bs + 1, dtype=torch.int32, pin_memory=True)
|
||||||
|
else:
|
||||||
|
seg_lens = generate_sequence_lengths(forward_batch, device="cpu")
|
||||||
|
indptr = torch.zeros(bs + 1, dtype=torch.int32, pin_memory=True)
|
||||||
|
torch.cumsum(seg_lens, dim=0, out=indptr[1:])
|
||||||
|
return indptr
|
||||||
|
|
||||||
def init_cuda_graph_batch_info(
|
def init_cuda_graph_batch_info(
|
||||||
self,
|
self,
|
||||||
max_bs_in_cuda_graph: int,
|
max_bs_in_cuda_graph: int,
|
||||||
@@ -224,6 +236,8 @@ class ChunkedSgmvLoRABackend(BaseLoRABackend):
|
|||||||
scalings=torch.zeros(self.max_loras_per_batch, dtype=torch.float),
|
scalings=torch.zeros(self.max_loras_per_batch, dtype=torch.float),
|
||||||
num_segments=None, # Set per batch
|
num_segments=None, # Set per batch
|
||||||
max_len=None, # Not used in CSGMV backend
|
max_len=None, # Not used in CSGMV backend
|
||||||
|
req_seg_indptr=torch.zeros(max_bs_in_cuda_graph + 1, dtype=torch.int32),
|
||||||
|
req_weight_indices=torch.zeros(max_bs_in_cuda_graph, dtype=torch.int32),
|
||||||
)
|
)
|
||||||
|
|
||||||
def prepare_lora_batch(
|
def prepare_lora_batch(
|
||||||
@@ -254,9 +268,15 @@ class ChunkedSgmvLoRABackend(BaseLoRABackend):
|
|||||||
scalings, dtype=torch.float, pin_memory=True, device="cpu"
|
scalings, dtype=torch.float, pin_memory=True, device="cpu"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
bs = forward_batch.batch_size
|
||||||
|
req_wi_tensor = torch.tensor(
|
||||||
|
weight_indices, dtype=torch.int32, pin_memory=True, device="cpu"
|
||||||
|
)
|
||||||
|
req_seg_indptr_cpu = self._build_req_seg_indptr(forward_batch)
|
||||||
|
|
||||||
if not use_cuda_graph:
|
if not use_cuda_graph:
|
||||||
batch_info = LoRABatchInfo(
|
batch_info = LoRABatchInfo(
|
||||||
bs=forward_batch.batch_size,
|
bs=bs,
|
||||||
num_segments=num_segments,
|
num_segments=num_segments,
|
||||||
max_len=chunk_size,
|
max_len=chunk_size,
|
||||||
use_cuda_graph=False,
|
use_cuda_graph=False,
|
||||||
@@ -275,12 +295,17 @@ class ChunkedSgmvLoRABackend(BaseLoRABackend):
|
|||||||
permutation=torch.empty(
|
permutation=torch.empty(
|
||||||
(len(permutation),), dtype=torch.int32, device=self.device
|
(len(permutation),), dtype=torch.int32, device=self.device
|
||||||
),
|
),
|
||||||
# Not used in chunked kernels
|
|
||||||
seg_lens=None,
|
seg_lens=None,
|
||||||
|
req_seg_indptr=torch.empty(
|
||||||
|
(bs + 1,), dtype=torch.int32, device=self.device
|
||||||
|
),
|
||||||
|
req_weight_indices=torch.empty(
|
||||||
|
(bs,), dtype=torch.int32, device=self.device
|
||||||
|
),
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
batch_info = self.cuda_graph_batch_info
|
batch_info = self.cuda_graph_batch_info
|
||||||
batch_info.bs = forward_batch.batch_size
|
batch_info.bs = bs
|
||||||
batch_info.num_segments = num_segments
|
batch_info.num_segments = num_segments
|
||||||
batch_info.max_len = chunk_size
|
batch_info.max_len = chunk_size
|
||||||
|
|
||||||
@@ -296,6 +321,8 @@ class ChunkedSgmvLoRABackend(BaseLoRABackend):
|
|||||||
)
|
)
|
||||||
batch_info.seg_indptr[: num_segments + 1].copy_(seg_indptr, non_blocking=True)
|
batch_info.seg_indptr[: num_segments + 1].copy_(seg_indptr, non_blocking=True)
|
||||||
batch_info.permutation[: len(permutation)].copy_(permutation, non_blocking=True)
|
batch_info.permutation[: len(permutation)].copy_(permutation, non_blocking=True)
|
||||||
|
batch_info.req_seg_indptr[: bs + 1].copy_(req_seg_indptr_cpu, non_blocking=True)
|
||||||
|
batch_info.req_weight_indices[:bs].copy_(req_wi_tensor, non_blocking=True)
|
||||||
|
|
||||||
self.batch_info = batch_info
|
self.batch_info = batch_info
|
||||||
self.lm_head_batch_info, self.lm_head_pass_batch_infos = (
|
self.lm_head_batch_info, self.lm_head_pass_batch_infos = (
|
||||||
|
|||||||
@@ -925,25 +925,37 @@ class FusedMoEWithLoRA(BaseLayerWithLoRA):
|
|||||||
max_lora_rank = self.down_lora_a_weights.shape[2]
|
max_lora_rank = self.down_lora_a_weights.shape[2]
|
||||||
|
|
||||||
cg_buffers = getattr(self.lora_backend, "moe_cg_buffers", None)
|
cg_buffers = getattr(self.lora_backend, "moe_cg_buffers", None)
|
||||||
|
wi = (
|
||||||
|
batch_info.req_weight_indices
|
||||||
|
if batch_info.req_weight_indices is not None
|
||||||
|
else batch_info.weight_indices
|
||||||
|
)
|
||||||
if cg_buffers is not None and batch_info.use_cuda_graph:
|
if cg_buffers is not None and batch_info.use_cuda_graph:
|
||||||
adapter_enabled = cg_buffers["adapter_enabled"]
|
adapter_enabled = cg_buffers["adapter_enabled"]
|
||||||
adapter_enabled.zero_()
|
adapter_enabled.zero_()
|
||||||
idx_buf = cg_buffers["weight_indices_long"]
|
idx_buf = cg_buffers["weight_indices_long"]
|
||||||
idx_buf[: batch_info.bs] = batch_info.weight_indices[: batch_info.bs]
|
idx_buf[: batch_info.bs] = wi[: batch_info.bs]
|
||||||
adapter_enabled.index_fill_(0, idx_buf[: batch_info.bs], 1)
|
adapter_enabled.index_fill_(0, idx_buf[: batch_info.bs], 1)
|
||||||
else:
|
else:
|
||||||
adapter_enabled = torch.zeros(
|
adapter_enabled = torch.zeros(
|
||||||
len(lora_ranks), dtype=torch.int32, device=lora_ranks.device
|
len(lora_ranks), dtype=torch.int32, device=lora_ranks.device
|
||||||
)
|
)
|
||||||
adapter_enabled.index_fill_(0, batch_info.weight_indices.long(), 1)
|
adapter_enabled.index_fill_(0, wi.long(), 1)
|
||||||
|
|
||||||
|
seg_indptr = (
|
||||||
|
batch_info.req_seg_indptr
|
||||||
|
if batch_info.req_seg_indptr is not None
|
||||||
|
else batch_info.seg_indptr
|
||||||
|
)
|
||||||
|
req_to_lora = wi
|
||||||
|
|
||||||
return LoRAInfo(
|
return LoRAInfo(
|
||||||
gate_up_lora_a_weights=self.gate_up_lora_a_weights,
|
gate_up_lora_a_weights=self.gate_up_lora_a_weights,
|
||||||
gate_up_lora_b_weights=self.gate_up_lora_b_weights,
|
gate_up_lora_b_weights=self.gate_up_lora_b_weights,
|
||||||
down_lora_a_weights=self.down_lora_a_weights,
|
down_lora_a_weights=self.down_lora_a_weights,
|
||||||
down_lora_b_weights=self.down_lora_b_weights,
|
down_lora_b_weights=self.down_lora_b_weights,
|
||||||
seg_indptr=batch_info.seg_indptr,
|
seg_indptr=seg_indptr,
|
||||||
req_to_lora=batch_info.weight_indices,
|
req_to_lora=req_to_lora,
|
||||||
lora_ranks=lora_ranks,
|
lora_ranks=lora_ranks,
|
||||||
adapter_enabled=adapter_enabled,
|
adapter_enabled=adapter_enabled,
|
||||||
max_lora_rank=max_lora_rank,
|
max_lora_rank=max_lora_rank,
|
||||||
|
|||||||
@@ -48,6 +48,16 @@ class LoRABatchInfo:
|
|||||||
# Computed from Python lists in prepare_lora_batch to avoid GPU sync.
|
# Computed from Python lists in prepare_lora_batch to avoid GPU sync.
|
||||||
has_active_lora: bool = False
|
has_active_lora: bool = False
|
||||||
|
|
||||||
|
# Per-request segment indptrs, shape (bs + 1,). Required by MoE virtual
|
||||||
|
# experts which map tokens to requests regardless of the dense-LoRA
|
||||||
|
# backend's internal segmentation. For the triton backend these are
|
||||||
|
# identical to seg_indptr/weight_indices; for csgmv they differ because
|
||||||
|
# its segments are chunked across adapters.
|
||||||
|
req_seg_indptr: Optional[torch.Tensor] = None
|
||||||
|
|
||||||
|
# Per-request adapter index, shape (bs,).
|
||||||
|
req_weight_indices: Optional[torch.Tensor] = None
|
||||||
|
|
||||||
|
|
||||||
class LoRAType(Enum):
|
class LoRAType(Enum):
|
||||||
LORA_A = 0
|
LORA_A = 0
|
||||||
|
|||||||
@@ -7055,10 +7055,6 @@ class ServerArgs:
|
|||||||
), "--max-lora-chunk-size must be a power of 2 between 16 and 128."
|
), "--max-lora-chunk-size must be a power of 2 between 16 and 128."
|
||||||
|
|
||||||
if self.lora_use_virtual_experts:
|
if self.lora_use_virtual_experts:
|
||||||
assert self.lora_backend == "triton", (
|
|
||||||
"--lora-use-virtual-experts requires --lora-backend triton. "
|
|
||||||
f"Got: {self.lora_backend}"
|
|
||||||
)
|
|
||||||
logger.info("Virtual expert computation enabled.")
|
logger.info("Virtual expert computation enabled.")
|
||||||
|
|
||||||
assert (
|
assert (
|
||||||
|
|||||||
Reference in New Issue
Block a user