[LoRA] Torch native backend: rework implementation and updated tests (#15187)

This commit is contained in:
Vladimir Serov
2025-12-30 11:48:51 +08:00
committed by GitHub
parent 0d003e34b0
commit 9e263c2162
14 changed files with 958 additions and 1000 deletions
@@ -5,7 +5,7 @@ from sglang.srt.lora.triton_ops import (
chunked_sgmv_lora_expand_forward, chunked_sgmv_lora_expand_forward,
chunked_sgmv_lora_shrink_forward, chunked_sgmv_lora_shrink_forward,
) )
from sglang.srt.lora.utils import LoRABatchInfo from sglang.srt.lora.utils import LoRABatchInfo, generate_sequence_lengths
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
@@ -283,24 +283,7 @@ class ChunkedSgmvLoRABackend(BaseLoRABackend):
""" """
with torch.device("cpu"): with torch.device("cpu"):
seq_weight_indices = torch.tensor(seq_weight_indices, dtype=torch.int32) seq_weight_indices = torch.tensor(seq_weight_indices, dtype=torch.int32)
seg_lens_cpu = generate_sequence_lengths(forward_batch)
if forward_batch.forward_mode.is_decode():
seg_lens_cpu = torch.ones(forward_batch.batch_size, dtype=torch.int32)
elif forward_batch.forward_mode.is_target_verify():
seg_lens_cpu = torch.full(
size=(forward_batch.batch_size,),
fill_value=forward_batch.spec_info.draft_token_num,
dtype=torch.int32,
)
elif forward_batch.forward_mode.is_extend():
seg_lens_cpu = torch.tensor(
forward_batch.extend_seq_lens_cpu,
dtype=torch.int32,
)
else:
raise ValueError(
f"Unsupported forward mode: {forward_batch.forward_mode}"
)
row_weight_indices = torch.repeat_interleave( row_weight_indices = torch.repeat_interleave(
seq_weight_indices, seg_lens_cpu seq_weight_indices, seg_lens_cpu
+127 -149
View File
@@ -1,11 +1,29 @@
from dataclasses import dataclass
from typing import Optional
import torch import torch
from sglang.srt.lora.backend.base_backend import BaseLoRABackend from sglang.srt.lora.backend.base_backend import BaseLoRABackend
from sglang.srt.lora.torch_ops import sgmv_expand, sgmv_expand_slice, sgmv_shrink from sglang.srt.lora.torch_ops import sgemm_lora_a_fwd, sgemm_lora_b_fwd
from sglang.srt.lora.utils import LoRABatchInfo from sglang.srt.lora.utils import LoRABatchInfo, generate_sequence_lengths
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
@dataclass
class TorchNativeLoRABatchInfo(LoRABatchInfo):
# ranks of each lora adapter, in shape (lora_num,) placed on cpu device
lora_ranks_cpu: Optional[torch.Tensor] = None
# Indice pointers of each segment in shape (num_segments + 1, ) placed on cpu device
seg_indptr_cpu: Optional[torch.Tensor] = None
# Lengths of each segments in shape (num_segments,) placed on cpu device
seg_lens_cpu: Optional[torch.Tensor] = None
# The index of lora adapter used by each segment, in shape (num_segments,) placed on cpu device
weight_indices_cpu: Optional[torch.Tensor] = None
class TorchNativeLoRABackend(BaseLoRABackend): class TorchNativeLoRABackend(BaseLoRABackend):
name = "torch_native" name = "torch_native"
@@ -20,27 +38,15 @@ class TorchNativeLoRABackend(BaseLoRABackend):
def run_lora_a_sgemm( def run_lora_a_sgemm(
self, x: torch.Tensor, weights: torch.Tensor, *args, **kwargs self, x: torch.Tensor, weights: torch.Tensor, *args, **kwargs
) -> torch.Tensor: ) -> torch.Tensor:
output_tensor = sgemm_lora_a_fwd(
total_seq_len, _ = x.shape inputs=x,
_, weight_out_dim, _ = weights.shape weights=weights,
weight_indices=self.batch_info.weight_indices_cpu,
output_tensor = torch.zeros( seg_len_tensor=self.batch_info.seg_lens_cpu,
(total_seq_len, weight_out_dim), dtype=x.dtype, device=x.device lora_ranks=self.batch_info.lora_ranks_cpu,
scaling_tensor=self.batch_info.scalings,
num_slices=1,
) )
sgmv_shrink(
x,
weights,
output_tensor,
self.batch_info.seg_lens,
self.batch_info.weight_indices,
1.0,
)
scaling = torch.repeat_interleave(
self.batch_info.scalings[self.batch_info.weight_indices],
self.batch_info.seg_lens,
output_size=total_seq_len,
).unsqueeze(-1)
output_tensor = output_tensor * scaling
return output_tensor return output_tensor
@@ -52,23 +58,18 @@ class TorchNativeLoRABackend(BaseLoRABackend):
*args, *args,
**kwargs, **kwargs,
) -> torch.Tensor: ) -> torch.Tensor:
total_seq_len, _ = x.shape
_, weight_out_dim, _ = weights.shape _, weight_out_dim, _ = weights.shape
output_offset = torch.tensor(
if base_output is None: [0, weight_out_dim], dtype=torch.int32, device="cpu"
output_tensor = torch.zeros( )
(total_seq_len, weight_out_dim), device=x.device, dtype=x.dtype output_tensor = sgemm_lora_b_fwd(
) inputs=x,
else: weights=weights,
output_tensor = base_output weight_indices=self.batch_info.weight_indices_cpu,
seg_len_tensor=self.batch_info.seg_lens_cpu,
sgmv_expand( lora_ranks=self.batch_info.lora_ranks_cpu,
x, slice_offsets=output_offset,
weights, base_output=base_output,
output_tensor,
self.batch_info.seg_lens,
self.batch_info.weight_indices,
True,
) )
return output_tensor return output_tensor
@@ -86,52 +87,25 @@ class TorchNativeLoRABackend(BaseLoRABackend):
**kwargs, **kwargs,
) -> torch.Tensor: ) -> torch.Tensor:
num_slices = 3 num_slices = 3
assert isinstance(qkv_lora_b, torch.Tensor) lora_a_output = sgemm_lora_a_fwd(
inputs=x,
total_seq_len, _ = x.shape weights=qkv_lora_a,
_, weight_intermediate_dim, _ = qkv_lora_a.shape weight_indices=self.batch_info.weight_indices_cpu,
_, weight_out_dim, _ = qkv_lora_b.shape seg_len_tensor=self.batch_info.seg_lens_cpu,
max_rank = weight_intermediate_dim // num_slices lora_ranks=self.batch_info.lora_ranks_cpu,
scaling_tensor=self.batch_info.scalings,
if base_output is None: num_slices=num_slices,
output_tensor = torch.zeros(
(total_seq_len, weight_out_dim), device=x.device, dtype=x.dtype
)
else:
output_tensor = base_output
lora_a_output = torch.zeros(
total_seq_len, weight_intermediate_dim, dtype=x.dtype, device=x.device
) )
sgmv_shrink(
x,
qkv_lora_a,
lora_a_output,
self.batch_info.seg_lens,
self.batch_info.weight_indices,
1.0,
)
scaling = torch.repeat_interleave(
self.batch_info.scalings[self.batch_info.weight_indices],
self.batch_info.seg_lens,
output_size=total_seq_len,
).unsqueeze(-1)
lora_a_output = lora_a_output * scaling
for slice_id in range(num_slices): output_tensor = sgemm_lora_b_fwd(
slice_offset = output_offset_cpu[slice_id] inputs=lora_a_output,
slice_offset_next = output_offset_cpu[slice_id + 1] weights=qkv_lora_b,
slice_size = slice_offset_next - slice_offset weight_indices=self.batch_info.weight_indices_cpu,
sgmv_expand_slice( seg_len_tensor=self.batch_info.seg_lens_cpu,
lora_a_output[:, (max_rank * slice_id) : (max_rank * (slice_id + 1))], lora_ranks=self.batch_info.lora_ranks_cpu,
qkv_lora_b[:, slice_offset:slice_offset_next], slice_offsets=output_offset_cpu,
output_tensor, base_output=base_output,
self.batch_info.seg_lens, )
self.batch_info.weight_indices,
slice_offset,
slice_size,
True,
)
return output_tensor return output_tensor
@@ -144,54 +118,32 @@ class TorchNativeLoRABackend(BaseLoRABackend):
*args, *args,
**kwargs, **kwargs,
) -> torch.Tensor: ) -> torch.Tensor:
num_slices = 2 num_slices = 2
assert isinstance(gate_up_lora_b, torch.Tensor)
total_seq_len, _ = x.shape
_, weight_intermediate_dim, _ = gate_up_lora_a.shape
_, weight_out_dim, _ = gate_up_lora_b.shape _, weight_out_dim, _ = gate_up_lora_b.shape
slice_size = weight_out_dim // num_slices slice_size = weight_out_dim // num_slices
max_rank = weight_intermediate_dim // num_slices output_offset = torch.tensor(
[0, slice_size, weight_out_dim], dtype=torch.int32, device="cpu"
if base_output is None:
output_tensor = torch.zeros(
(total_seq_len, weight_out_dim), device=x.device, dtype=x.dtype
)
else:
output_tensor = base_output
lora_a_output = torch.zeros(
total_seq_len, weight_intermediate_dim, dtype=x.dtype, device=x.device
) )
sgmv_shrink(
x,
gate_up_lora_a,
lora_a_output,
self.batch_info.seg_lens,
self.batch_info.weight_indices,
1.0,
)
scaling = torch.repeat_interleave(
self.batch_info.scalings[self.batch_info.weight_indices],
self.batch_info.seg_lens,
output_size=total_seq_len,
).unsqueeze(-1)
lora_a_output = lora_a_output * scaling
slice_offset = 0 lora_a_output = sgemm_lora_a_fwd(
for slice_id in range(num_slices): inputs=x,
sgmv_expand_slice( weights=gate_up_lora_a,
lora_a_output[:, (max_rank * slice_id) : (max_rank * (slice_id + 1))], weight_indices=self.batch_info.weight_indices_cpu,
gate_up_lora_b[:, slice_offset : slice_offset + slice_size], seg_len_tensor=self.batch_info.seg_lens_cpu,
output_tensor, lora_ranks=self.batch_info.lora_ranks_cpu,
self.batch_info.seg_lens, scaling_tensor=self.batch_info.scalings,
self.batch_info.weight_indices, num_slices=num_slices,
slice_offset, )
slice_size,
True, output_tensor = sgemm_lora_b_fwd(
) inputs=lora_a_output,
slice_offset += slice_size weights=gate_up_lora_b,
weight_indices=self.batch_info.weight_indices_cpu,
seg_len_tensor=self.batch_info.seg_lens_cpu,
lora_ranks=self.batch_info.lora_ranks_cpu,
slice_offsets=output_offset,
base_output=base_output,
)
return output_tensor return output_tensor
@@ -201,19 +153,19 @@ class TorchNativeLoRABackend(BaseLoRABackend):
num_tokens_per_bs: int, num_tokens_per_bs: int,
): ):
with torch.device("cuda"): with torch.device("cuda"):
self.cuda_graph_batch_info = LoRABatchInfo( self.cuda_graph_batch_info = TorchNativeLoRABatchInfo(
bs=max_bs_in_cuda_graph,
use_cuda_graph=True, use_cuda_graph=True,
num_segments=None, bs=max_bs_in_cuda_graph,
num_segments=self.max_loras_per_batch,
seg_lens=torch.full( seg_lens=torch.full(
(max_bs_in_cuda_graph,), num_tokens_per_bs, dtype=torch.int32 (max_bs_in_cuda_graph,), num_tokens_per_bs, dtype=torch.int32
), ),
seg_indptr=torch.empty(max_bs_in_cuda_graph + 1, dtype=torch.int32), seg_indptr=torch.zeros(max_bs_in_cuda_graph + 1, dtype=torch.int32),
max_len=num_tokens_per_bs,
weight_indices=torch.zeros(max_bs_in_cuda_graph, dtype=torch.int32), weight_indices=torch.zeros(max_bs_in_cuda_graph, dtype=torch.int32),
lora_ranks=torch.zeros(self.max_loras_per_batch, dtype=torch.int32), lora_ranks=torch.zeros(self.max_loras_per_batch, dtype=torch.int32),
scalings=torch.zeros(self.max_loras_per_batch, dtype=torch.float), scalings=torch.zeros(self.max_loras_per_batch, dtype=torch.float),
permutation=None, permutation=None,
max_len=num_tokens_per_bs,
) )
# Initialize seg_indptr for CUDA graph as they remain constant # Initialize seg_indptr for CUDA graph as they remain constant
@@ -232,10 +184,36 @@ class TorchNativeLoRABackend(BaseLoRABackend):
scalings: list[float], scalings: list[float],
use_cuda_graph: bool, use_cuda_graph: bool,
): ):
# Use pinned memory to avoid synchronizations during host-to-device transfer original_seq_lens_cpu = generate_sequence_lengths(forward_batch, device="cpu")
weight_indices_tensor = torch.tensor( original_weight_indices_tensor = torch.tensor(
weight_indices, dtype=torch.int32, pin_memory=True, device="cpu" weight_indices, dtype=torch.int32, device="cpu"
) )
unique_weight_indices_tensor, inverse_weight_indices_tensor = (
torch.unique_consecutive(
original_weight_indices_tensor, return_inverse=True
)
)
seg_lens_cpu = (
torch.zeros_like(
unique_weight_indices_tensor, dtype=torch.int32, device="cpu"
)
.scatter_add_(
0,
inverse_weight_indices_tensor,
original_seq_lens_cpu,
)
.pin_memory()
)
seg_indptr_cpu = torch.zeros(
(len(seg_lens_cpu) + 1,), dtype=torch.int32, pin_memory=True
)
seg_indptr_cpu[1:] = torch.cumsum(seg_lens_cpu, dim=0)
# Use pinned memory to avoid synchronizations during host-to-device transfer
weight_indices_tensor = unique_weight_indices_tensor.pin_memory()
lora_ranks_tensor = torch.tensor( lora_ranks_tensor = torch.tensor(
lora_ranks, dtype=torch.int32, pin_memory=True, device="cpu" lora_ranks, dtype=torch.int32, pin_memory=True, device="cpu"
) )
@@ -253,27 +231,17 @@ class TorchNativeLoRABackend(BaseLoRABackend):
batch_info.bs = forward_batch.batch_size batch_info.bs = forward_batch.batch_size
batch_info.num_segments = forward_batch.batch_size batch_info.num_segments = forward_batch.batch_size
else: else:
max_len = ( max_len = max(seg_lens_cpu)
# Calculate max_len from the CPU copy to avoid D2H transfer.
max(forward_batch.extend_seq_lens_cpu)
if forward_batch.forward_mode.is_extend()
else 1
)
seg_lens = (
forward_batch.extend_seq_lens
if forward_batch.forward_mode.is_extend()
else torch.ones(bs, dtype=torch.int32, device=self.device)
)
seg_indptr = torch.zeros((bs + 1,), dtype=torch.int32, device=self.device)
seg_indptr[1:] = torch.cumsum(seg_lens, dim=0)
batch_info = LoRABatchInfo( batch_info = TorchNativeLoRABatchInfo(
bs=forward_batch.batch_size, bs=forward_batch.batch_size,
num_segments=forward_batch.batch_size, num_segments=forward_batch.batch_size,
max_len=max_len, max_len=max_len,
use_cuda_graph=False, use_cuda_graph=False,
seg_lens=seg_lens, seg_lens=torch.empty((bs,), dtype=torch.int32, device=self.device),
seg_indptr=seg_indptr, seg_indptr=torch.empty(
(bs + 1,), dtype=torch.int32, device=self.device
),
weight_indices=torch.empty( weight_indices=torch.empty(
(bs,), dtype=torch.int32, device=self.device (bs,), dtype=torch.int32, device=self.device
), ),
@@ -294,4 +262,14 @@ class TorchNativeLoRABackend(BaseLoRABackend):
scalings_tensor, non_blocking=True scalings_tensor, non_blocking=True
) )
batch_info.weight_indices[:bs].copy_(weight_indices_tensor, non_blocking=True) batch_info.weight_indices[:bs].copy_(weight_indices_tensor, non_blocking=True)
batch_info.seg_indptr[: len(seg_indptr_cpu)].copy_(
seg_indptr_cpu, non_blocking=True
)
batch_info.seg_lens[: len(seg_lens_cpu)].copy_(seg_lens_cpu, non_blocking=True)
batch_info.lora_ranks_cpu = lora_ranks_tensor
batch_info.seg_indptr_cpu = seg_indptr_cpu
batch_info.seg_lens_cpu = seg_lens_cpu
batch_info.weight_indices_cpu = weight_indices_tensor
self.batch_info = batch_info self.batch_info = batch_info
+3 -4
View File
@@ -1,7 +1,6 @@
from .lora_ops import sgmv_expand, sgmv_expand_slice, sgmv_shrink from .lora_ops import sgemm_lora_a_fwd, sgemm_lora_b_fwd
__all__ = [ __all__ = [
"sgmv_expand", "sgemm_lora_a_fwd",
"sgmv_expand_slice", "sgemm_lora_b_fwd",
"sgmv_shrink",
] ]
+85 -101
View File
@@ -1,125 +1,109 @@
# SPDX-License-Identifier: Apache-2.0 from typing import Optional
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import torch import torch
def sgmv_expand( def sgemm_lora_a_fwd(
inputs: torch.Tensor, inputs: torch.Tensor,
lora_b_weights: torch.Tensor, weights: torch.Tensor,
output_tensor: torch.Tensor, weight_indices: torch.Tensor,
seq_len_tensor: torch.Tensor, seg_len_tensor: torch.Tensor,
lora_indices_tensor: torch.Tensor, lora_ranks: torch.Tensor,
add_inputs: bool = False, scaling_tensor: torch.Tensor,
num_slices: int = 1,
): ):
total_seq_len, _ = inputs.shape total_seq_len, input_dim = inputs.shape
exploded_indices = torch.repeat_interleave( if weights.numel() == 0:
lora_indices_tensor, seq_len_tensor, output_size=total_seq_len return torch.zeros(total_seq_len, 0, dtype=inputs.dtype, device=inputs.device)
num_loras, weight_out_dim, _ = weights.shape
max_rank = weight_out_dim // num_slices
output = torch.zeros(
total_seq_len, num_slices * max_rank, dtype=inputs.dtype, device=inputs.device
) )
bgmv_expand(inputs, lora_b_weights, output_tensor, exploded_indices, add_inputs) token_offset = 0
for lora_idx, seq_len, rank in zip(
weight_indices, seg_len_tensor, lora_ranks[weight_indices]
):
if seq_len == 0:
continue
if rank > 0:
x_seq = inputs[token_offset : token_offset + seq_len, :]
w_seq = weights[lora_idx, : num_slices * rank, :]
result = torch.einsum("si, oi -> so", x_seq, w_seq)
output[token_offset : token_offset + seq_len, : num_slices * rank] = (
scaling_tensor[lora_idx] * result
)
token_offset += seq_len
return output
def bgmv_expand( def sgemm_lora_b_fwd(
inputs: torch.Tensor, inputs: torch.Tensor,
lora_b_weights: torch.Tensor, weights: torch.Tensor,
output_tensor: torch.Tensor, weight_indices: torch.Tensor,
lora_indices_tensor: torch.Tensor, seg_len_tensor: torch.Tensor,
add_inputs: bool = True, lora_ranks: torch.Tensor,
slice_offsets: torch.Tensor,
base_output: Optional[torch.Tensor] = None,
): ):
selected_loras = lora_b_weights[lora_indices_tensor].to(dtype=output_tensor.dtype) total_seq_len, _ = inputs.shape
if len(selected_loras.shape) == 4: num_loras, weight_out_dim, _ = weights.shape
selected_loras = selected_loras.squeeze(dim=1) total_output_dim = slice_offsets[-1].item() if len(slice_offsets) > 0 else 0
inputs = inputs.to(dtype=output_tensor.dtype)
outputs = torch.einsum("bi, boi -> bo", inputs, selected_loras)
limit = output_tensor.shape[0] if weights.numel() == 0:
if outputs.shape[0] == 1 and output_tensor.shape[0] != 1: return torch.zeros(
limit = 1 total_seq_len, total_output_dim, dtype=inputs.dtype, device=inputs.device
)
# LoRA adapter and model may add different amounts of padding to output num_slices = len(slice_offsets) - 1
common_len = min(outputs.shape[1], output_tensor.shape[1])
if add_inputs: if base_output is not None:
output_tensor[:, :common_len] += outputs[:limit, :common_len] output = base_output
else: else:
output_tensor[:, :common_len] = outputs[:limit, :common_len] output = torch.zeros(
total_seq_len, total_output_dim, dtype=inputs.dtype, device=inputs.device
)
token_offset = 0
for lora_idx, seq_len, rank in zip(
weight_indices, seg_len_tensor, lora_ranks[weight_indices]
):
if seq_len == 0:
continue
def sgmv_shrink( if rank == 0:
inputs: torch.Tensor, token_offset += seq_len
lora_a_weights: torch.Tensor, continue
output_tensor: torch.Tensor,
seq_len_tensor: torch.Tensor,
lora_indices_tensor: torch.Tensor,
scaling: float,
):
total_seq_len, _ = inputs.shape
exploded_indices = torch.repeat_interleave(
lora_indices_tensor, seq_len_tensor, output_size=total_seq_len
)
bgmv_shrink(inputs, lora_a_weights, output_tensor, exploded_indices, scaling) for slice_idx in range(num_slices):
slice_start_input = slice_idx * rank
slice_end_input = (slice_idx + 1) * rank
slice_start_output = slice_offsets[slice_idx]
slice_end_output = slice_offsets[slice_idx + 1]
def bgmv_shrink( x_slice = inputs[
inputs: torch.Tensor, token_offset : token_offset + seq_len :,
lora_a_weights: torch.Tensor, slice_start_input:slice_end_input,
output_tensor: torch.Tensor, ] # (seq_len, rank)
lora_indices_tensor: torch.Tensor, w_slice = weights[
scaling: float = 1.0, lora_idx, slice_start_output:slice_end_output, :rank
): ] # (slice_dim, rank)
selected_loras = lora_a_weights[lora_indices_tensor].to(dtype=output_tensor.dtype)
if len(selected_loras.shape) == 4:
selected_loras = selected_loras.squeeze(dim=1)
inputs = inputs.to(dtype=output_tensor.dtype)
outputs = torch.einsum("bi, boi -> bo", inputs, selected_loras)
output_tensor[:, : outputs.shape[1]] = scaling * outputs[:] result = torch.einsum("si, oi -> so", x_slice, w_slice)
output[
token_offset : token_offset + seq_len,
slice_start_output:slice_end_output,
] += result
token_offset += seq_len
def sgmv_expand_slice( return output
inputs: torch.Tensor,
lora_b_weights: torch.Tensor,
output_tensor: torch.Tensor,
seq_len_tensor: torch.Tensor,
lora_indices_tensor: torch.Tensor,
slice_offset: int,
slice_size: int,
add_inputs: bool = False,
):
total_seq_len, _ = inputs.shape
exploded_indices = torch.repeat_interleave(
lora_indices_tensor, seq_len_tensor, output_size=total_seq_len
)
bgmv_expand_slice(
inputs,
lora_b_weights,
output_tensor,
exploded_indices,
slice_offset,
slice_size,
add_inputs,
)
def bgmv_expand_slice(
inputs: torch.Tensor,
lora_b_weights: torch.Tensor,
output_tensor: torch.Tensor,
lora_indices_tensor: torch.Tensor,
slice_offset: int,
slice_size: int,
add_inputs: bool = True,
):
selected_loras = lora_b_weights[lora_indices_tensor].to(dtype=output_tensor.dtype)
inputs = inputs.to(dtype=output_tensor.dtype)
if len(selected_loras.shape) == 4:
selected_loras = selected_loras.squeeze(dim=1)
outputs = torch.einsum("bi, boi -> bo", inputs, selected_loras)
if add_inputs:
output_tensor[:, slice_offset : slice_offset + slice_size] += outputs[:]
else:
output_tensor[:, slice_offset : slice_offset + slice_size] = outputs[:]
+29
View File
@@ -4,6 +4,7 @@ from typing import Iterable, Optional, Set, Tuple
import torch import torch
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.utils.hf_transformers_utils import AutoConfig from sglang.srt.utils.hf_transformers_utils import AutoConfig
@@ -151,3 +152,31 @@ def get_target_module_name(full_module_name: str, target_modules: Set[str]) -> s
EMBEDDING_NAMES = ["embed_tokens", "lm_head"] EMBEDDING_NAMES = ["embed_tokens", "lm_head"]
ROW_PARALLELISM_LINEAR_LORA_NAMES = ["o_proj", "down_proj"] ROW_PARALLELISM_LINEAR_LORA_NAMES = ["o_proj", "down_proj"]
def generate_sequence_lengths(
forward_batch: ForwardBatch, device: Optional[torch.device] = None
) -> torch.Tensor:
device = torch.get_default_device() if device is None else device
with torch.device(device):
if forward_batch.forward_mode.is_decode():
seg_lens = torch.ones(forward_batch.batch_size, dtype=torch.int32)
elif forward_batch.forward_mode.is_target_verify():
seg_lens = torch.full(
size=(forward_batch.batch_size,),
fill_value=forward_batch.spec_info.draft_token_num,
dtype=torch.int32,
)
elif forward_batch.forward_mode.is_extend():
seg_lens = (
forward_batch.extend_seq_lens
if forward_batch.extend_seq_lens.device == device
else torch.tensor(
forward_batch.extend_seq_lens_cpu,
dtype=torch.int32,
)
)
else:
raise ValueError(f"Unsupported forward mode: {forward_batch.forward_mode}")
return seg_lens
+146 -1
View File
@@ -1,6 +1,6 @@
import dataclasses import dataclasses
import random import random
from typing import List from typing import List, Optional
import torch import torch
@@ -116,6 +116,151 @@ ALL_OTHER_MULTI_LORA_MODELS = [
] ]
def safe_matmul(a: torch.Tensor, b: torch.Tensor) -> torch.Tensor:
"""Matrix multiplication with mixed precision handling for float16"""
result = torch.matmul(a.float(), b.float())
return result.to(a.dtype)
def reference_sgmv_shrink(
x: torch.Tensor,
weights: torch.Tensor,
weight_indices: torch.Tensor,
seq_lengths: torch.Tensor,
lora_ranks: torch.Tensor,
lora_scalings: torch.Tensor,
num_slices: int = 1,
) -> torch.Tensor:
"""
Simple sequence-level reference implementation of SGMV shrink operation.
Args:
x: (total_seq_len, input_dim) - Input activations
weights: (num_loras, num_slices * max_rank, input_dim) - LoRA A weights
weight_indices: LoRA idx for each sequence
seq_lengths: Length of each sequence
lora_ranks: LoRA rank for each LoRA adapters
lora_scalings: LoRA scaling for each LoRA adapters
num_slices: Number of slices (3 for QKV, 2 for gate_up, 1 for others)
Returns:
output: (total_seq_len, num_slices * max_rank) - Intermediate activations
"""
if weights.numel() == 0:
total_seq_len = x.shape[0]
return torch.zeros(total_seq_len, 0, dtype=x.dtype, device=x.device)
total_seq_len, _ = x.shape
_, weight_out_dim, _ = weights.shape
max_rank = weight_out_dim // num_slices
output = torch.zeros(
total_seq_len, num_slices * max_rank, dtype=x.dtype, device=x.device
)
token_offset = 0
for lora_idx, seq_len, rank, scaling in zip(
weight_indices,
seq_lengths,
lora_ranks[weight_indices],
lora_scalings[weight_indices],
):
if seq_len == 0:
continue
if rank > 0:
x_seq = x[token_offset : token_offset + seq_len, :]
w_seq = weights[lora_idx, : num_slices * rank, :]
result = safe_matmul(x_seq, w_seq.t())
output[token_offset : token_offset + seq_len, : num_slices * rank] = (
scaling * result
)
token_offset += seq_len
return output
def reference_sgmv_expand(
x: torch.Tensor,
weights: torch.Tensor,
weight_indices: torch.Tensor,
seq_lengths: torch.Tensor,
lora_ranks: torch.Tensor,
slice_offsets: torch.Tensor,
base_output: Optional[torch.Tensor] = None,
) -> torch.Tensor:
"""
Simple sequence-level reference implementation of SGMV expand operation.
Args:
x: (total_seq_len, num_slices * max_rank) - Intermediate activations
weights: (num_loras, output_dim, max_rank) - LoRA B weights
weight_indices: LoRA idx for each sequence
seq_lengths: Length of each sequence
lora_ranks: LoRA rank for each LoRA adapters
slice_offsets: Tensor defining slice boundaries
base_output: Optional base output to accumulate into
Returns:
output: (total_seq_len, total_output_dim) - Final output
"""
if weights.numel() == 0:
total_seq_len = x.shape[0]
total_output_dim = slice_offsets[-1].item() if len(slice_offsets) > 0 else 0
return torch.zeros(
total_seq_len, total_output_dim, dtype=x.dtype, device=x.device
)
total_seq_len, _ = x.shape
num_slices = len(slice_offsets) - 1
if base_output is not None:
output = base_output.clone()
else:
total_output_dim = slice_offsets[-1].item()
output = torch.zeros(
total_seq_len, total_output_dim, dtype=x.dtype, device=x.device
)
token_offset = 0
for lora_idx, seq_len, rank in zip(
weight_indices, seq_lengths, lora_ranks[weight_indices]
):
if seq_len == 0:
continue
if rank > 0:
# Extract sequence intermediate activations
x_seq = x[
token_offset : token_offset + seq_len, : num_slices * rank
] # (seq_len, num_slices * rank)
for slice_idx in range(num_slices):
slice_start_input = slice_idx * rank
slice_end_input = (slice_idx + 1) * rank
slice_start_output = slice_offsets[slice_idx].item()
slice_end_output = slice_offsets[slice_idx + 1].item()
x_slice = x_seq[:, slice_start_input:slice_end_input] # (seq_len, rank)
w_slice = weights[
lora_idx, slice_start_output:slice_end_output, :rank
] # (slice_dim, rank)
result = safe_matmul(x_slice, w_slice.t()) # (seq_len, slice_dim)
output[
token_offset : token_offset + seq_len,
slice_start_output:slice_end_output,
] += result
token_offset += seq_len
return output
def run_lora_test_one_by_one( def run_lora_test_one_by_one(
prompts: List[str], prompts: List[str],
model_case: LoRAModelCase, model_case: LoRAModelCase,
+96 -209
View File
@@ -1,7 +1,7 @@
import random import random
import unittest import unittest
from enum import Enum from enum import Enum
from typing import Dict, List, Optional, Tuple from typing import List, Optional, Tuple
import torch import torch
@@ -13,6 +13,7 @@ from sglang.srt.lora.triton_ops import (
from sglang.srt.lora.triton_ops.chunked_sgmv_expand import _chunked_lora_expand_kernel from sglang.srt.lora.triton_ops.chunked_sgmv_expand import _chunked_lora_expand_kernel
from sglang.srt.lora.triton_ops.chunked_sgmv_shrink import _chunked_lora_shrink_kernel from sglang.srt.lora.triton_ops.chunked_sgmv_shrink import _chunked_lora_shrink_kernel
from sglang.srt.lora.utils import LoRABatchInfo from sglang.srt.lora.utils import LoRABatchInfo
from sglang.test.lora_utils import reference_sgmv_expand, reference_sgmv_shrink
CHUNK_SIZE = 16 CHUNK_SIZE = 16
@@ -22,12 +23,6 @@ def reset_kernel_cache():
_chunked_lora_expand_kernel._clear_cache() _chunked_lora_expand_kernel._clear_cache()
def safe_matmul(a: torch.Tensor, b: torch.Tensor) -> torch.Tensor:
"""Matrix multiplication with mixed precision handling for float16"""
result = torch.matmul(a.float(), b.float())
return result.to(a.dtype)
class BatchComposition(Enum): class BatchComposition(Enum):
UNIFORM = "uniform" UNIFORM = "uniform"
MIXED = "mixed" MIXED = "mixed"
@@ -38,150 +33,7 @@ class BatchComposition(Enum):
class BatchMode(Enum): class BatchMode(Enum):
PREFILL = "prefill" PREFILL = "prefill"
DECODE = "decode" DECODE = "decode"
TARGET_VERIFY = "verify"
def reference_sgmv_shrink(
x: torch.Tensor,
weights: torch.Tensor,
batch_info: LoRABatchInfo,
seq_lengths: List[int],
lora_assignments: List[str],
num_slices: int = 1,
) -> torch.Tensor:
"""
Simple sequence-level reference implementation of SGMV shrink operation.
Args:
x: (total_seq_len, input_dim) - Input activations
weights: (num_loras, num_slices * max_rank, input_dim) - LoRA A weights
batch_info: Batch information (only used for lora_ranks)
seq_lengths: Length of each sequence
lora_assignments: LoRA name for each sequence
num_slices: Number of slices (3 for QKV, 2 for gate_up, 1 for others)
Returns:
output: (total_seq_len, num_slices * max_rank) - Intermediate activations
"""
if weights.numel() == 0:
total_seq_len = x.shape[0]
return torch.zeros(total_seq_len, 0, dtype=x.dtype, device=x.device)
total_seq_len, input_dim = x.shape
num_loras, weight_out_dim, _ = weights.shape
max_rank = weight_out_dim // num_slices
output = torch.zeros(
total_seq_len, num_slices * max_rank, dtype=x.dtype, device=x.device
)
unique_loras = sorted(set(lora_assignments))
lora_name_to_idx = {name: idx for idx, name in enumerate(unique_loras)}
lora_ranks = batch_info.lora_ranks.cpu().numpy()
token_offset = 0
for seq_len, lora_name in zip(seq_lengths, lora_assignments):
if seq_len == 0:
continue
lora_idx = lora_name_to_idx[lora_name]
rank = lora_ranks[lora_idx]
if rank > 0:
x_seq = x[token_offset : token_offset + seq_len, :]
w_seq = weights[lora_idx, : num_slices * rank, :]
result = safe_matmul(x_seq, w_seq.t())
output[token_offset : token_offset + seq_len, : num_slices * rank] = result
token_offset += seq_len
return output
def reference_sgmv_expand(
x: torch.Tensor,
weights: torch.Tensor,
batch_info: LoRABatchInfo,
seq_lengths: List[int],
lora_assignments: List[str],
slice_offsets: torch.Tensor,
max_slice_size: int,
base_output: Optional[torch.Tensor] = None,
) -> torch.Tensor:
"""
Simple sequence-level reference implementation of SGMV expand operation.
Args:
x: (total_seq_len, num_slices * max_rank) - Intermediate activations
weights: (num_loras, output_dim, max_rank) - LoRA B weights
batch_info: Batch information (only used for lora_ranks)
seq_lengths: Length of each sequence
lora_assignments: LoRA name for each sequence
slice_offsets: Tensor defining slice boundaries
max_slice_size: Maximum slice size for chunking
base_output: Optional base output to accumulate into
Returns:
output: (total_seq_len, total_output_dim) - Final output
"""
if weights.numel() == 0:
total_seq_len = x.shape[0]
total_output_dim = slice_offsets[-1].item() if len(slice_offsets) > 0 else 0
return torch.zeros(
total_seq_len, total_output_dim, dtype=x.dtype, device=x.device
)
total_seq_len, _ = x.shape
num_slices = len(slice_offsets) - 1
if base_output is not None:
output = base_output.clone()
else:
total_output_dim = slice_offsets[-1].item()
output = torch.zeros(
total_seq_len, total_output_dim, dtype=x.dtype, device=x.device
)
unique_loras = sorted(set(lora_assignments))
lora_name_to_idx = {name: idx for idx, name in enumerate(unique_loras)}
lora_ranks = batch_info.lora_ranks.cpu().numpy()
token_offset = 0
for seq_len, lora_name in zip(seq_lengths, lora_assignments):
if seq_len == 0:
continue
lora_idx = lora_name_to_idx[lora_name]
lora_rank = lora_ranks[lora_idx]
if lora_rank > 0:
# Extract sequence intermediate activations
x_seq = x[
token_offset : token_offset + seq_len, : num_slices * lora_rank
] # (seq_len, num_slices * rank)
for slice_idx in range(num_slices):
slice_start_input = slice_idx * lora_rank
slice_end_input = (slice_idx + 1) * lora_rank
slice_start_output = slice_offsets[slice_idx].item()
slice_end_output = slice_offsets[slice_idx + 1].item()
x_slice = x_seq[:, slice_start_input:slice_end_input] # (seq_len, rank)
w_slice = weights[
lora_idx, slice_start_output:slice_end_output, :lora_rank
] # (slice_dim, rank)
result = safe_matmul(x_slice, w_slice.t()) # (seq_len, slice_dim)
output[
token_offset : token_offset + seq_len,
slice_start_output:slice_end_output,
] += result
token_offset += seq_len
return output
class TestChunkedSGMV(unittest.TestCase): class TestChunkedSGMV(unittest.TestCase):
@@ -196,7 +48,7 @@ class TestChunkedSGMV(unittest.TestCase):
chunked_output: torch.Tensor, chunked_output: torch.Tensor,
reference_output: torch.Tensor, reference_output: torch.Tensor,
seq_lengths: List[int], seq_lengths: List[int],
lora_assignments: List[str], lora_assignments: List[int],
batch_info: LoRABatchInfo, batch_info: LoRABatchInfo,
num_slices: int, num_slices: int,
test_name: str, test_name: str,
@@ -207,19 +59,15 @@ class TestChunkedSGMV(unittest.TestCase):
The chunked SGMV shrink kernel only guarantees correctness for The chunked SGMV shrink kernel only guarantees correctness for
output[seq_start:seq_end, :rank * num_slices] for each sequence. output[seq_start:seq_end, :rank * num_slices] for each sequence.
""" """
# Create mapping from LoRA names to indices and ranks
unique_loras = sorted(set(lora_assignments))
lora_name_to_idx = {name: idx for idx, name in enumerate(unique_loras)}
lora_ranks = batch_info.lora_ranks.cpu().numpy() lora_ranks = batch_info.lora_ranks.cpu().numpy()
token_offset = 0 token_offset = 0
for seq_idx, (seq_len, lora_name) in enumerate( for seq_idx, (lora_idx, seq_len) in enumerate(
zip(seq_lengths, lora_assignments) zip(lora_assignments, seq_lengths)
): ):
if seq_len == 0: if seq_len == 0:
continue continue
lora_idx = lora_name_to_idx[lora_name]
rank = lora_ranks[lora_idx] rank = lora_ranks[lora_idx]
if rank > 0: if rank > 0:
@@ -238,7 +86,7 @@ class TestChunkedSGMV(unittest.TestCase):
reference_seq, reference_seq,
rtol=self.RTOL, rtol=self.RTOL,
atol=self.ATOL, atol=self.ATOL,
msg=f"Shrink operation failed for {test_name}, sequence {seq_idx} ({lora_name})", msg=f"Shrink operation failed for {test_name}, sequence {seq_idx} ({lora_idx})",
) )
token_offset += seq_len token_offset += seq_len
@@ -318,23 +166,22 @@ class TestChunkedSGMV(unittest.TestCase):
def create_batch_info( def create_batch_info(
self, self,
lora_names: List[str],
seq_lengths: List[int], seq_lengths: List[int],
lora_assignments: List[Optional[str]], lora_assignments: List[Optional[int]],
batch_mode: BatchMode = BatchMode.PREFILL, batch_mode: BatchMode = BatchMode.PREFILL,
) -> LoRABatchInfo: ) -> LoRABatchInfo:
"""Create LoRABatchInfo using the same logic as chunked backend""" """Create LoRABatchInfo using the same logic as chunked backend"""
unique_loras = sorted(set(lora_assignments)) lora_ranks = [self.lora_configs[name][0] for name in lora_names]
lora_name_to_idx = {name: idx for idx, name in enumerate(unique_loras)}
seq_weight_indices = [lora_name_to_idx[name] for name in lora_assignments]
lora_ranks = [self.lora_configs[name][0] for name in unique_loras]
def create_mock_batch(): def create_mock_batch():
# Create a minimal mock ForwardBatch for the test # Create a minimal mock ForwardBatch for the test
class MockForwardBatch: class MockForwardBatch:
def __init__(self, batch_size, seq_lengths): def __init__(self, batch_size, seq_lengths, device):
self.batch_size = batch_size self.batch_size = batch_size
self.extend_seq_lens = torch.tensor(
seq_lengths, dtype=torch.int32, device=device
)
self.extend_seq_lens_cpu = seq_lengths self.extend_seq_lens_cpu = seq_lengths
self.forward_mode = MockForwardMode() self.forward_mode = MockForwardMode()
@@ -342,13 +189,22 @@ class TestChunkedSGMV(unittest.TestCase):
def is_extend(self): def is_extend(self):
return batch_mode == BatchMode.PREFILL return batch_mode == BatchMode.PREFILL
return MockForwardBatch(len(seq_lengths), seq_lengths) def is_decode(self):
return batch_mode == BatchMode.DECODE
def is_target_verify(self):
return batch_mode == BatchMode.TARGET_VERIFY
def is_prefill(self):
return self.is_extend()
return MockForwardBatch(len(seq_lengths), seq_lengths, self.device)
mock_batch = create_mock_batch() mock_batch = create_mock_batch()
# Use the same functions as chunked backend # Use the same functions as chunked backend
permutation, weights_reordered = ChunkedSgmvLoRABackend._get_permutation( permutation, weights_reordered = ChunkedSgmvLoRABackend._get_permutation(
seq_weight_indices, mock_batch lora_assignments, mock_batch
) )
# Create a minimal backend instance to access _get_segments_info # Create a minimal backend instance to access _get_segments_info
@@ -363,7 +219,7 @@ class TestChunkedSGMV(unittest.TestCase):
chunk_size=CHUNK_SIZE, chunk_size=CHUNK_SIZE,
) )
scalings = [1.0] * len(unique_loras) scalings = [1.0] * len(lora_names)
seg_indptr_tensor = seg_indptr.to(self.device) seg_indptr_tensor = seg_indptr.to(self.device)
weight_indices_tensor = weight_indices_list.to(self.device) weight_indices_tensor = weight_indices_list.to(self.device)
lora_ranks_tensor = ( lora_ranks_tensor = (
@@ -437,7 +293,7 @@ class TestChunkedSGMV(unittest.TestCase):
include_missing_k: bool = False, include_missing_k: bool = False,
) -> Tuple[ ) -> Tuple[
torch.Tensor, torch.Tensor,
Dict[str, Tuple[torch.Tensor, torch.Tensor]], List[Tuple[torch.Tensor, torch.Tensor]],
LoRABatchInfo, LoRABatchInfo,
List[int], List[int],
List[str], List[str],
@@ -451,20 +307,21 @@ class TestChunkedSGMV(unittest.TestCase):
batch_size, batch_mode, 1, self.max_seq_len batch_size, batch_mode, 1, self.max_seq_len
) )
if batch_composition == BatchComposition.UNIFORM: if batch_composition == BatchComposition.UNIFORM:
lora_assignments = ["lora_A"] * batch_size lora_names = ["lora_A"]
lora_assignments = [lora_names.index("lora_A")] * batch_size
elif batch_composition == BatchComposition.MIXED: elif batch_composition == BatchComposition.MIXED:
lora_names = ["lora_A", "lora_B", "lora_C", None] lora_names = ["lora_A", "lora_B", "lora_C", None]
lora_assignments = [ lora_assignments = [(i % len(lora_names)) for i in range(batch_size)]
lora_names[i % len(lora_names)] for i in range(batch_size)
]
elif batch_composition == BatchComposition.SKEWED: elif batch_composition == BatchComposition.SKEWED:
lora_names = ["lora_A", "lora_B"]
num_minority = max(1, batch_size // 8) num_minority = max(1, batch_size // 8)
lora_assignments = ["lora_A"] * num_minority + ["lora_B"] * ( lora_assignments = [lora_names.index("lora_A")] * num_minority + [
batch_size - num_minority lora_names.index("lora_B")
) ] * (batch_size - num_minority)
random.shuffle(lora_assignments) random.shuffle(lora_assignments)
elif batch_composition == BatchComposition.NONE: elif batch_composition == BatchComposition.NONE:
lora_assignments = [None] * batch_size lora_names = [None]
lora_assignments = [0] * batch_size
else: else:
raise ValueError(f"Unknown batch composition: {batch_composition}") raise ValueError(f"Unknown batch composition: {batch_composition}")
@@ -473,39 +330,45 @@ class TestChunkedSGMV(unittest.TestCase):
total_seq_len, self.input_dim, dtype=self.dtype, device=self.device total_seq_len, self.input_dim, dtype=self.dtype, device=self.device
) )
normalized_assignments = [ normalized_lora_names = [
name if name is not None else "_NO_LORA_" for name in lora_assignments "_NO_LORA_" if name is None else name for name in lora_names
] ]
unique_loras = set(normalized_assignments) weights = []
weights = {} for lora_name in normalized_lora_names:
for lora_name in unique_loras: weights.append(self.create_lora_weights(lora_name, include_missing_k))
weights[lora_name] = self.create_lora_weights(lora_name, include_missing_k)
batch_info = self.create_batch_info( batch_info = self.create_batch_info(
seq_lengths, normalized_assignments, batch_mode normalized_lora_names, seq_lengths, lora_assignments, batch_mode
) )
return x, weights, batch_info, seq_lengths, normalized_assignments return x, weights, batch_info, seq_lengths, lora_assignments
def run_test_comparison( def run_test_comparison(
self, self,
x: torch.Tensor, x: torch.Tensor,
weights: Dict[str, Tuple[torch.Tensor, torch.Tensor]], weights: List[Tuple[torch.Tensor, torch.Tensor]],
batch_info: LoRABatchInfo, batch_info: LoRABatchInfo,
seq_lengths: List[int], seq_lengths: List[int],
lora_assignments: List[str], lora_assignments: List[int],
test_name: str, test_name: str,
): ):
"""Run comparison between chunked and reference implementations""" """Run comparison between chunked and reference implementations"""
if not weights: # Handle case with no LoRA weights if not weights: # Handle case with no LoRA weights
return return
lora_assignments_tensor = torch.tensor(
lora_assignments, dtype=torch.int32, device="cpu"
)
seq_lengths_tensor = torch.tensor(seq_lengths, dtype=torch.int32, device="cpu")
lora_ranks_tensor = batch_info.lora_ranks.detach().cpu()
scalings_tensor = batch_info.scalings.detach().cpu()
# Stack LoRA A weights # Stack LoRA A weights
lora_a_weights = [weights[name][0] for name in sorted(weights.keys())] lora_a_weights = [weight[0] for weight in weights]
stacked_lora_a = self.stack_lora_weights(lora_a_weights, is_lora_a=True) stacked_lora_a = self.stack_lora_weights(lora_a_weights, is_lora_a=True)
# Stack LoRA B weights # Stack LoRA B weights
lora_b_weights = [weights[name][1] for name in sorted(weights.keys())] lora_b_weights = [weight[1] for weight in weights]
stacked_lora_b = self.stack_lora_weights(lora_b_weights, is_lora_a=False) stacked_lora_b = self.stack_lora_weights(lora_b_weights, is_lora_a=False)
# Test shrink operation # Test shrink operation
@@ -513,7 +376,13 @@ class TestChunkedSGMV(unittest.TestCase):
x, stacked_lora_a, batch_info, num_slices=3 x, stacked_lora_a, batch_info, num_slices=3
) )
reference_shrink = reference_sgmv_shrink( reference_shrink = reference_sgmv_shrink(
x, stacked_lora_a, batch_info, seq_lengths, lora_assignments, num_slices=3 x,
stacked_lora_a,
lora_assignments_tensor,
seq_lengths_tensor,
lora_ranks_tensor,
scalings_tensor,
num_slices=3,
) )
# Only compare valid portions of shrink output (first rank * num_slices columns per sequence) # Only compare valid portions of shrink output (first rank * num_slices columns per sequence)
@@ -539,11 +408,10 @@ class TestChunkedSGMV(unittest.TestCase):
reference_expand = reference_sgmv_expand( reference_expand = reference_sgmv_expand(
reference_shrink, reference_shrink,
stacked_lora_b, stacked_lora_b,
batch_info, lora_assignments_tensor,
seq_lengths, seq_lengths_tensor,
lora_assignments, lora_ranks_tensor,
self.slice_offsets, self.slice_offsets,
self.max_slice_size,
) )
torch.testing.assert_close( torch.testing.assert_close(
@@ -564,7 +432,16 @@ class TestChunkedSGMV(unittest.TestCase):
self.create_test_batch(BatchComposition.UNIFORM, batch_size) self.create_test_batch(BatchComposition.UNIFORM, batch_size)
) )
lora_a_weights = [weights[name][0] for name in sorted(weights.keys())] lora_assignments_tensor = torch.tensor(
lora_assignments, dtype=torch.int32, device="cpu"
)
seq_lengths_tensor = torch.tensor(
seq_lengths, dtype=torch.int32, device="cpu"
)
lora_ranks_tensor = batch_info.lora_ranks.detach().cpu()
scalings_tensor = batch_info.scalings.detach().cpu()
lora_a_weights = [weight[0] for weight in weights]
stacked_lora_a = self.stack_lora_weights(lora_a_weights, is_lora_a=True) stacked_lora_a = self.stack_lora_weights(lora_a_weights, is_lora_a=True)
chunked_shrink = chunked_sgmv_lora_shrink_forward( chunked_shrink = chunked_sgmv_lora_shrink_forward(
@@ -573,9 +450,10 @@ class TestChunkedSGMV(unittest.TestCase):
reference_shrink = reference_sgmv_shrink( reference_shrink = reference_sgmv_shrink(
x, x,
stacked_lora_a, stacked_lora_a,
batch_info, lora_assignments_tensor,
seq_lengths, seq_lengths_tensor,
lora_assignments, lora_ranks_tensor,
scalings_tensor,
num_slices=3, num_slices=3,
) )
@@ -591,19 +469,29 @@ class TestChunkedSGMV(unittest.TestCase):
self.create_test_batch(BatchComposition.UNIFORM, batch_size) self.create_test_batch(BatchComposition.UNIFORM, batch_size)
) )
lora_a_weights = [weights[name][0] for name in sorted(weights.keys())] lora_assignments_tensor = torch.tensor(
lora_assignments, dtype=torch.int32, device="cpu"
)
seq_lengths_tensor = torch.tensor(
seq_lengths, dtype=torch.int32, device="cpu"
)
lora_ranks_tensor = batch_info.lora_ranks.detach().cpu()
scalings_tensor = batch_info.scalings.detach().cpu()
lora_a_weights = [weight[0] for weight in weights]
stacked_lora_a = self.stack_lora_weights(lora_a_weights, is_lora_a=True) stacked_lora_a = self.stack_lora_weights(lora_a_weights, is_lora_a=True)
intermediate = reference_sgmv_shrink( intermediate = reference_sgmv_shrink(
x, x,
stacked_lora_a, stacked_lora_a,
batch_info, lora_assignments_tensor,
seq_lengths, seq_lengths_tensor,
lora_assignments, lora_ranks_tensor,
scalings_tensor,
num_slices=3, num_slices=3,
) )
lora_b_weights = [weights[name][1] for name in sorted(weights.keys())] lora_b_weights = [weight[1] for weight in weights]
stacked_lora_b = self.stack_lora_weights( stacked_lora_b = self.stack_lora_weights(
lora_b_weights, is_lora_a=False lora_b_weights, is_lora_a=False
) )
@@ -619,11 +507,10 @@ class TestChunkedSGMV(unittest.TestCase):
reference_expand = reference_sgmv_expand( reference_expand = reference_sgmv_expand(
intermediate, intermediate,
stacked_lora_b, stacked_lora_b,
batch_info, lora_assignments_tensor,
seq_lengths, seq_lengths_tensor,
lora_assignments, lora_ranks_tensor,
self.slice_offsets, self.slice_offsets,
self.max_slice_size,
) )
torch.testing.assert_close( torch.testing.assert_close(
+1 -2
View File
@@ -17,7 +17,7 @@ import os
import unittest import unittest
from typing import List from typing import List
from lora_utils import ( from sglang.test.lora_utils import (
ALL_OTHER_LORA_MODELS, ALL_OTHER_LORA_MODELS,
CI_LORA_MODELS, CI_LORA_MODELS,
DEFAULT_PROMPTS, DEFAULT_PROMPTS,
@@ -26,7 +26,6 @@ from lora_utils import (
run_lora_test_by_batch, run_lora_test_by_batch,
run_lora_test_one_by_one, run_lora_test_one_by_one,
) )
from sglang.test.test_utils import CustomTestCase, is_in_ci from sglang.test.test_utils import CustomTestCase, is_in_ci
TEST_CUDA_GRAPH_PADDING_PROMPTS = [ TEST_CUDA_GRAPH_PADDING_PROMPTS = [
+218
View File
@@ -0,0 +1,218 @@
import random
import unittest
import torch
from sglang.srt.lora.torch_ops.lora_ops import sgemm_lora_a_fwd, sgemm_lora_b_fwd
from sglang.test.lora_utils import reference_sgmv_expand, reference_sgmv_shrink
from sglang.test.test_utils import CustomTestCase
class TestLoraOps(CustomTestCase):
def test_sgemm_lora_a_fwd(self):
batch_size = 2
input_dim = 1024
num_loras = 3
dtype = torch.float32
possible_lora_ranks = [8, 16, 32, 64, 128, 256]
lora_ranks = random.sample(
possible_lora_ranks,
counts=[num_loras] * len(possible_lora_ranks),
k=num_loras,
)
max_lora_rank = max(lora_ranks)
possible_lora_scaling = [0.25, 0.5, 1.0, 2.0, 4.0]
lora_scaling = random.sample(
possible_lora_scaling,
counts=[num_loras] * len(possible_lora_scaling),
k=num_loras,
)
inputs = torch.randn(batch_size, input_dim, dtype=dtype)
lora_a_weights = torch.randn(num_loras, max_lora_rank, input_dim, dtype=dtype)
lora_indices_tensor = torch.randint(
num_loras, (batch_size,), dtype=torch.int32, device="cpu"
)
seq_len_tensor = torch.ones(batch_size, dtype=torch.int32, device="cpu")
lora_ranks_tensor = torch.tensor(lora_ranks, dtype=torch.int32, device="cpu")
lora_scaling_tensor = torch.tensor(
lora_scaling, dtype=torch.float16, device="cpu"
)
expect_output = reference_sgmv_shrink(
inputs,
lora_a_weights,
lora_indices_tensor,
seq_len_tensor,
lora_ranks_tensor,
lora_scaling_tensor,
)
actual_output = sgemm_lora_a_fwd(
inputs,
lora_a_weights,
lora_indices_tensor,
seq_len_tensor,
lora_ranks_tensor,
lora_scaling_tensor,
)
self.assertTrue(torch.allclose(actual_output, expect_output))
def test_sgemm_lora_b_fwd(self):
batch_size = 2
output_dim = 1024
num_loras = 3
dtype = torch.float32
possible_lora_ranks = [8, 16, 32, 64, 128, 256]
lora_ranks = random.sample(
possible_lora_ranks,
counts=[num_loras] * len(possible_lora_ranks),
k=num_loras,
)
max_lora_rank = max(lora_ranks)
inputs = torch.randn(batch_size, max_lora_rank, dtype=dtype)
lora_b_weights = torch.randn(num_loras, output_dim, max_lora_rank, dtype=dtype)
lora_ranks_tensor = torch.tensor(lora_ranks, dtype=torch.int32, device="cpu")
seq_len_tensor = torch.ones(batch_size, dtype=torch.int32, device="cpu")
lora_indices_tensor = torch.randint(
num_loras, (batch_size,), dtype=torch.int32, device="cpu"
)
slice_offsets = torch.tensor([0, output_dim], dtype=torch.int32, device="cpu")
expect_output = reference_sgmv_expand(
inputs,
lora_b_weights,
lora_indices_tensor,
seq_len_tensor,
lora_ranks_tensor,
slice_offsets,
)
actual_output = sgemm_lora_b_fwd(
inputs,
lora_b_weights,
lora_indices_tensor,
seq_len_tensor,
lora_ranks_tensor,
slice_offsets,
)
self.assertTrue(torch.allclose(actual_output, expect_output))
def test_sgemm_lora_a_fwd_expand(self):
batch_size = 2
input_dim = 1024
num_loras = 3
dtype = torch.float32
possible_lora_ranks = [8, 16, 32, 64, 128, 256]
lora_ranks = random.sample(
possible_lora_ranks,
counts=[num_loras] * len(possible_lora_ranks),
k=num_loras,
)
max_lora_rank = max(lora_ranks)
possible_lora_scaling = [0.25, 0.5, 1.0, 2.0, 4.0]
lora_scaling = random.sample(
possible_lora_scaling,
counts=[num_loras] * len(possible_lora_scaling),
k=num_loras,
)
seq_len_tensor = torch.randint(
num_loras, (batch_size,), dtype=torch.int32, device="cpu"
)
seq_len = sum(seq_len_tensor)
inputs = torch.randn(seq_len, input_dim, dtype=dtype)
lora_a_weights = torch.randn(num_loras, max_lora_rank, input_dim, dtype=dtype)
lora_indices_tensor = torch.randint(
num_loras, (batch_size,), dtype=torch.int32, device="cpu"
)
lora_ranks_tensor = torch.tensor(lora_ranks, dtype=torch.int32, device="cpu")
lora_scaling_tensor = torch.tensor(
lora_scaling, dtype=torch.float16, device="cpu"
)
expect_output = reference_sgmv_shrink(
inputs,
lora_a_weights,
lora_indices_tensor,
seq_len_tensor,
lora_ranks_tensor,
lora_scaling_tensor,
)
actual_output = sgemm_lora_a_fwd(
inputs,
lora_a_weights,
lora_indices_tensor,
seq_len_tensor,
lora_ranks_tensor,
lora_scaling_tensor,
)
self.assertTrue(torch.allclose(actual_output, expect_output))
def test_sgemm_lora_b_fwd_expand(self):
batch_size = 2
output_dim = 1024
num_loras = 3
dtype = torch.float32
possible_lora_ranks = [8, 16, 32, 64, 128, 256]
lora_ranks = random.sample(
possible_lora_ranks,
counts=[num_loras] * len(possible_lora_ranks),
k=num_loras,
)
max_lora_rank = max(lora_ranks)
seq_len_tensor = torch.randint(
num_loras, (batch_size,), dtype=torch.int32, device="cpu"
)
seq_len = sum(seq_len_tensor)
inputs = torch.randn(seq_len, max_lora_rank, dtype=dtype)
lora_b_weights = torch.randn(num_loras, output_dim, max_lora_rank, dtype=dtype)
lora_ranks_tensor = torch.tensor(lora_ranks, dtype=torch.int32, device="cpu")
lora_indices_tensor = torch.randint(
num_loras, (batch_size,), dtype=torch.int32, device="cpu"
)
slice_offsets = torch.tensor([0, output_dim], dtype=torch.int32, device="cpu")
expect_output = reference_sgmv_expand(
inputs,
lora_b_weights,
lora_indices_tensor,
seq_len_tensor,
lora_ranks_tensor,
slice_offsets,
)
actual_output = sgemm_lora_b_fwd(
inputs,
lora_b_weights,
lora_indices_tensor,
seq_len_tensor,
lora_ranks_tensor,
slice_offsets,
)
self.assertTrue(torch.allclose(actual_output, expect_output))
if __name__ == "__main__":
unittest.main()
+6 -2
View File
@@ -2,10 +2,14 @@ import random
import unittest import unittest
from typing import Sequence from typing import Sequence
from lora_utils import TORCH_DTYPES, LoRAAdaptor, LoRAModelCase, ensure_reproducibility
from sglang.srt.models.qwen3_vl import Qwen3VLForConditionalGeneration from sglang.srt.models.qwen3_vl import Qwen3VLForConditionalGeneration
from sglang.srt.models.qwen3_vl_moe import Qwen3VLMoeForConditionalGeneration from sglang.srt.models.qwen3_vl_moe import Qwen3VLMoeForConditionalGeneration
from sglang.test.lora_utils import (
TORCH_DTYPES,
LoRAAdaptor,
LoRAModelCase,
ensure_reproducibility,
)
from sglang.test.runners import HFRunner, SRTRunner from sglang.test.runners import HFRunner, SRTRunner
from sglang.test.test_utils import CustomTestCase, calculate_rouge_l from sglang.test.test_utils import CustomTestCase, calculate_rouge_l
+1 -2
View File
@@ -15,13 +15,12 @@
import multiprocessing as mp import multiprocessing as mp
import unittest import unittest
from lora_utils import ( from sglang.test.lora_utils import (
CI_MULTI_LORA_MODELS, CI_MULTI_LORA_MODELS,
LoRAAdaptor, LoRAAdaptor,
LoRAModelCase, LoRAModelCase,
run_lora_multiple_batch_on_model_cases, run_lora_multiple_batch_on_model_cases,
) )
from sglang.test.test_utils import CustomTestCase from sglang.test.test_utils import CustomTestCase
LORA_MODELS_QWEN3 = [ LORA_MODELS_QWEN3 = [
+244
View File
@@ -0,0 +1,244 @@
import unittest
import torch
from sglang.srt.lora.backend.torch_backend import TorchNativeLoRABackend
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.test.lora_utils import reference_sgmv_expand, reference_sgmv_shrink
from sglang.test.test_utils import CustomTestCase
class TestTorchNativeLoRABackend(CustomTestCase):
device = "cpu"
weight_indices = [0, 1]
lora_ranks = [1, 1]
scalings = [1.0, 0.5]
seq_lens = [1, 1]
use_cuda_graph = False
forward_batch = ForwardBatch(
forward_mode=ForwardMode.EXTEND,
batch_size=2,
input_ids=torch.tensor([[1, 2, 3], [4, 5, 6]], dtype=torch.int32),
req_pool_indices=None,
seq_lens=None,
out_cache_loc=None,
seq_lens_sum=6,
extend_seq_lens=torch.tensor(seq_lens, dtype=torch.int32),
extend_seq_lens_cpu=seq_lens,
)
@classmethod
def setUpClass(cls):
cls.backend = TorchNativeLoRABackend(max_loras_per_batch=2, device=cls.device)
cls.backend.prepare_lora_batch(
forward_batch=cls.forward_batch,
weight_indices=cls.weight_indices,
lora_ranks=cls.lora_ranks,
scalings=cls.scalings,
use_cuda_graph=cls.use_cuda_graph,
)
def test_run_lora_a_sgemm(self):
batch_size = 2
input_dim = 4
output_dim = 6
num_loras = 3
dtype = torch.float32
x = torch.randn(batch_size, input_dim, dtype=dtype)
weights = torch.randn(num_loras, output_dim, input_dim, dtype=dtype)
weight_indices_tensor = torch.tensor(
self.weight_indices, dtype=torch.int32, device=self.device
)
seg_len_tensor = torch.tensor(
self.seq_lens, dtype=torch.int32, device=self.device
)
lora_ranks_tensor = torch.tensor(
self.lora_ranks, dtype=torch.int32, device=self.device
)
scalings_tensor = torch.tensor(
self.scalings, dtype=torch.float, device=self.device
)
expect_output = reference_sgmv_shrink(
x,
weights,
weight_indices_tensor,
seg_len_tensor,
lora_ranks_tensor,
scalings_tensor,
)
actual_output = self.backend.run_lora_a_sgemm(x, weights)
self.assertTrue(torch.allclose(actual_output, expect_output))
def test_run_lora_b_sgemm(self):
batch_size = 2
input_dim = 6
output_dim = 4
num_loras = 3
dtype = torch.float32
x = torch.randn(batch_size, input_dim, dtype=dtype)
weights = torch.randn(num_loras, output_dim, input_dim, dtype=dtype)
_, weight_out_dim, _ = weights.shape
weight_indices_tensor = torch.tensor(
self.weight_indices, dtype=torch.int32, device=self.device
)
seg_len_tensor = torch.tensor(
self.seq_lens, dtype=torch.int32, device=self.device
)
lora_ranks_tensor = torch.tensor(
self.lora_ranks, dtype=torch.int32, device=self.device
)
expect_output = reference_sgmv_expand(
x,
weights,
weight_indices_tensor,
seg_len_tensor,
lora_ranks_tensor,
slice_offsets=torch.tensor(
[0, weight_out_dim], dtype=torch.int32, device="cpu"
),
)
actual_output = self.backend.run_lora_b_sgemm(x, weights)
self.assertTrue(torch.allclose(actual_output, expect_output))
def test_run_qkv_lora(self):
batch_size = 2
num_loras = 3
input_dim = 6
output_offset = [0, 3, 6, 9, 12]
output_dim = output_offset[-1]
num_slices = len(output_offset) - 1
max_lora_rank = max(self.lora_ranks)
dtype = torch.float32
x = torch.randn(batch_size, input_dim, dtype=dtype)
output_offset_cpu = torch.tensor(output_offset, dtype=torch.int32)
qkv_lora_a = torch.randn(
num_loras, max_lora_rank * num_slices, input_dim, dtype=dtype
)
qkv_lora_b = torch.randn(
num_loras, output_dim, max_lora_rank * num_slices, dtype=dtype
)
weight_indices_tensor = torch.tensor(
self.weight_indices, dtype=torch.int32, device=self.device
)
seg_len_tensor = torch.tensor(
self.seq_lens, dtype=torch.int32, device=self.device
)
lora_ranks_tensor = torch.tensor(
self.lora_ranks, dtype=torch.int32, device=self.device
)
scalings_tensor = torch.tensor(
self.scalings, dtype=torch.float, device=self.device
)
expect_lora_a_output = reference_sgmv_shrink(
x,
qkv_lora_a,
weight_indices_tensor,
seg_len_tensor,
lora_ranks_tensor,
scalings_tensor,
num_slices,
)
expect_output = reference_sgmv_expand(
expect_lora_a_output,
qkv_lora_b,
weight_indices_tensor,
seg_len_tensor,
lora_ranks_tensor,
output_offset_cpu,
)
actual_output = self.backend.run_qkv_lora(
x, qkv_lora_a, qkv_lora_b, None, output_offset_cpu, 0
)
self.assertTrue(torch.allclose(actual_output, expect_output))
def test_run_gate_up_lora(self):
batch_size = 2
input_dim = 6
output_dim = 4
num_loras = 3
dtype = torch.float32
max_lora_rank = max(self.lora_ranks)
num_slices = 2
x = torch.randn(batch_size, input_dim, dtype=dtype)
gate_up_lora_a = torch.randn(
num_loras, max_lora_rank * num_slices, input_dim, dtype=dtype
)
gate_up_lora_b = torch.randn(
num_loras, output_dim, max_lora_rank * num_slices, dtype=dtype
)
_, weight_out_dim, _ = gate_up_lora_b.shape
slice_size = weight_out_dim // num_slices
output_offset = torch.tensor(
[0, slice_size, weight_out_dim], dtype=torch.int32, device="cpu"
)
weight_indices_tensor = torch.tensor(
self.weight_indices, dtype=torch.int32, device=self.device
)
seg_len_tensor = torch.tensor(
self.seq_lens, dtype=torch.int32, device=self.device
)
lora_ranks_tensor = torch.tensor(
self.lora_ranks, dtype=torch.int32, device=self.device
)
scalings_tensor = torch.tensor(
self.scalings, dtype=torch.float, device=self.device
)
expect_lora_a_output = reference_sgmv_shrink(
x,
gate_up_lora_a,
weight_indices_tensor,
seg_len_tensor,
lora_ranks_tensor,
scalings_tensor,
num_slices,
)
expect_output = reference_sgmv_expand(
expect_lora_a_output,
gate_up_lora_b,
weight_indices_tensor,
seg_len_tensor,
lora_ranks_tensor,
slice_offsets=output_offset,
)
actual_output = self.backend.run_gate_up_lora(x, gate_up_lora_a, gate_up_lora_b)
self.assertTrue(torch.allclose(actual_output, expect_output))
if __name__ == "__main__":
unittest.main()
-287
View File
@@ -1,287 +0,0 @@
import unittest
import torch
from sglang.srt.lora.torch_ops.lora_ops import (
bgmv_expand,
bgmv_expand_slice,
bgmv_shrink,
sgmv_expand,
sgmv_expand_slice,
sgmv_shrink,
)
from sglang.test.test_utils import CustomTestCase
class TestLoraOps(CustomTestCase):
def test_sgmv_expand(self):
batch_size = 2
input_dim = 4
output_dim = 6
num_loras = 3
dtype = torch.float32
inputs = torch.randn(batch_size, input_dim, dtype=dtype)
lora_b_weights = torch.randn(num_loras, output_dim, input_dim, dtype=dtype)
seq_len_tensor = torch.ones(batch_size, dtype=torch.int32)
lora_indices_tensor = torch.randint(0, num_loras, (batch_size,))
add_inputs = True
total_seq_len, _ = inputs.shape
exploded_indices = torch.repeat_interleave(
lora_indices_tensor, seq_len_tensor, output_size=total_seq_len
)
expect_output = torch.zeros(batch_size, output_dim, dtype=dtype)
bgmv_expand(inputs, lora_b_weights, expect_output, exploded_indices, add_inputs)
actual_output = torch.zeros(batch_size, output_dim, dtype=dtype)
sgmv_expand(
inputs,
lora_b_weights,
actual_output,
seq_len_tensor,
lora_indices_tensor,
add_inputs,
)
self.assertTrue(torch.allclose(actual_output, expect_output))
def test_bgmv_expand(self):
batch_size = 2
input_dim = 4
output_dim = 6
num_loras = 3
dtype = torch.float32
inputs = torch.randn(batch_size, input_dim, dtype=dtype)
lora_b_weights = torch.randn(num_loras, output_dim, input_dim, dtype=dtype)
lora_indices_tensor = torch.randint(0, num_loras, (batch_size,))
selected_loras = lora_b_weights[lora_indices_tensor].to(dtype=dtype)
selected_loras = selected_loras.squeeze(dim=1)
inputs = inputs.to(dtype=dtype)
outputs = torch.einsum("bi, boi -> bo", inputs, selected_loras)
limit = batch_size
common_len = min(outputs.shape[1], output_dim)
expect_output = torch.zeros(batch_size, output_dim, dtype=dtype)
expect_output[:, :common_len] = outputs[:limit, :common_len]
actual_output = torch.zeros(batch_size, output_dim, dtype=dtype)
bgmv_expand(
inputs,
lora_b_weights,
actual_output,
lora_indices_tensor,
add_inputs=False,
)
self.assertTrue(torch.allclose(actual_output, expect_output))
def test_bgmv_expand_add_residual(self):
batch_size = 2
input_dim = 4
output_dim = 6
num_loras = 3
dtype = torch.float32
inputs = torch.randn(batch_size, input_dim, dtype=dtype)
lora_b_weights = torch.randn(num_loras, output_dim, input_dim, dtype=dtype)
lora_indices_tensor = torch.randint(0, num_loras, (batch_size,))
selected_loras = lora_b_weights[lora_indices_tensor].to(dtype=dtype)
selected_loras = selected_loras.squeeze(dim=1)
inputs = inputs.to(dtype=dtype)
outputs = torch.einsum("bi, boi -> bo", inputs, selected_loras)
limit = batch_size
common_len = min(outputs.shape[1], output_dim)
expect_output = torch.randn(batch_size, output_dim, dtype=dtype)
actual_output = expect_output.clone()
expect_output[:, :common_len] += outputs[:limit, :common_len]
bgmv_expand(
inputs,
lora_b_weights,
actual_output,
lora_indices_tensor,
add_inputs=True,
)
self.assertTrue(torch.allclose(actual_output, expect_output))
def test_sgmv_shrink(self):
batch_size = 2
input_dim = 4
output_dim = 6
num_loras = 3
dtype = torch.float32
inputs = torch.randn(batch_size, input_dim, dtype=dtype)
lora_a_weights = torch.randn(num_loras, output_dim, input_dim, dtype=dtype)
seq_len_tensor = torch.ones(batch_size, dtype=torch.int32)
lora_indices_tensor = torch.randint(0, num_loras, (batch_size,))
scaling = 0.9
total_seq_len, _ = inputs.shape
exploded_indices = torch.repeat_interleave(
lora_indices_tensor, seq_len_tensor, output_size=total_seq_len
)
expect_output = torch.zeros(batch_size, output_dim, dtype=dtype)
bgmv_shrink(inputs, lora_a_weights, expect_output, exploded_indices, scaling)
actual_output = torch.zeros(batch_size, output_dim, dtype=dtype)
sgmv_shrink(
inputs,
lora_a_weights,
actual_output,
seq_len_tensor,
lora_indices_tensor,
scaling,
)
self.assertTrue(torch.allclose(actual_output, expect_output))
def test_bgmv_shrink(self):
batch_size = 2
input_dim = 4
output_dim = 6
num_loras = 3
dtype = torch.float32
inputs = torch.randn(batch_size, input_dim, dtype=dtype)
lora_a_weights = torch.randn(num_loras, output_dim, input_dim, dtype=dtype)
lora_indices_tensor = torch.randint(0, num_loras, (batch_size,))
scaling = 0.9
selected_loras = lora_a_weights[lora_indices_tensor].to(dtype=dtype)
inputs = inputs.to(dtype=dtype)
outputs = torch.einsum("bi, boi -> bo", inputs, selected_loras)
expect_output = torch.zeros(batch_size, output_dim, dtype=dtype)
expect_output[:, : outputs.shape[1]] = scaling * outputs[:]
actual_output = torch.zeros(batch_size, output_dim, dtype=dtype)
bgmv_shrink(
inputs,
lora_a_weights,
actual_output,
lora_indices_tensor,
scaling=scaling,
)
self.assertTrue(torch.allclose(actual_output, expect_output))
def test_sgmv_expand_slice(self):
batch_size = 2
input_dim = 4
output_dim = 6
output_dim_slice = 12
num_loras = 3
dtype = torch.float32
inputs = torch.randn(batch_size, input_dim, dtype=dtype)
lora_b_weights = torch.randn(num_loras, output_dim, input_dim, dtype=dtype)
seq_len_tensor = torch.ones(batch_size, dtype=torch.int32)
lora_indices_tensor = torch.randint(0, num_loras, (batch_size,))
slice_offset = 2
slice_size = 6
add_inputs = False
total_seq_len, _ = inputs.shape
exploded_indices = torch.repeat_interleave(
lora_indices_tensor, seq_len_tensor, output_size=total_seq_len
)
expect_output = torch.randn(batch_size, output_dim_slice, dtype=dtype)
actual_output = expect_output.clone()
bgmv_expand_slice(
inputs,
lora_b_weights,
expect_output,
exploded_indices,
slice_offset,
slice_size,
add_inputs,
)
sgmv_expand_slice(
inputs,
lora_b_weights,
actual_output,
seq_len_tensor,
lora_indices_tensor,
slice_offset,
slice_size,
add_inputs,
)
self.assertTrue(torch.allclose(actual_output, expect_output))
def test_bgmv_expand_slice(self):
batch_size = 2
input_dim = 4
output_dim = 6
output_dim_slice = 12
num_loras = 3
dtype = torch.float32
inputs = torch.randn(batch_size, input_dim, dtype=dtype)
lora_b_weights = torch.randn(num_loras, output_dim, input_dim, dtype=dtype)
lora_indices_tensor = torch.randint(0, num_loras, (batch_size,))
slice_offset = 2
slice_size = 6
selected_loras = lora_b_weights[lora_indices_tensor].to(dtype=dtype)
inputs = inputs.to(dtype=dtype)
outputs = torch.einsum("bi, boi -> bo", inputs, selected_loras)
expect_output = torch.zeros(batch_size, output_dim_slice, dtype=dtype)
expect_output[:, slice_offset : slice_offset + slice_size] = outputs[:]
actual_output = torch.zeros(batch_size, output_dim_slice, dtype=dtype)
bgmv_expand_slice(
inputs,
lora_b_weights,
actual_output,
lora_indices_tensor,
slice_offset,
slice_size,
add_inputs=False,
)
self.assertTrue(torch.allclose(actual_output, expect_output))
def test_bgmv_expand_slice_add_residual(self):
batch_size = 2
input_dim = 4
output_dim = 6
output_dim_slice = 12
num_loras = 3
dtype = torch.float32
inputs = torch.randn(batch_size, input_dim, dtype=dtype)
lora_b_weights = torch.randn(num_loras, output_dim, input_dim, dtype=dtype)
lora_indices_tensor = torch.randint(0, num_loras, (batch_size,))
slice_offset = 2
slice_size = 6
selected_loras = lora_b_weights[lora_indices_tensor].to(dtype=dtype)
inputs = inputs.to(dtype=dtype)
outputs = torch.einsum("bi, boi -> bo", inputs, selected_loras)
expect_output = torch.randn(batch_size, output_dim_slice, dtype=dtype)
actual_output = expect_output.clone()
expect_output[:, slice_offset : slice_offset + slice_size] += outputs[:]
bgmv_expand_slice(
inputs,
lora_b_weights,
actual_output,
lora_indices_tensor,
slice_offset,
slice_size,
add_inputs=True,
)
self.assertTrue(torch.allclose(actual_output, expect_output))
if __name__ == "__main__":
unittest.main()
-224
View File
@@ -1,224 +0,0 @@
import unittest
import torch
from sglang.srt.lora.backend.torch_backend import TorchNativeLoRABackend
from sglang.srt.lora.torch_ops.lora_ops import (
sgmv_expand,
sgmv_expand_slice,
sgmv_shrink,
)
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.test.test_utils import CustomTestCase
class TestTorchNativeLoRABackend(CustomTestCase):
device = "cpu"
forward_batch = ForwardBatch(
forward_mode=ForwardMode.EXTEND,
batch_size=2,
input_ids=torch.tensor([[1, 2, 3], [4, 5, 6]], dtype=torch.int32),
req_pool_indices=None,
seq_lens=None,
out_cache_loc=None,
seq_lens_sum=6,
extend_seq_lens=torch.tensor([1, 1], dtype=torch.int32),
extend_seq_lens_cpu=[1, 1],
)
weight_indices = [0, 1]
lora_ranks = [1, 1]
scalings = [1.0, 0.5]
use_cuda_graph = False
@classmethod
def setUpClass(cls):
cls.backend = TorchNativeLoRABackend(max_loras_per_batch=2, device=cls.device)
cls.backend.prepare_lora_batch(
forward_batch=cls.forward_batch,
weight_indices=cls.weight_indices,
lora_ranks=cls.lora_ranks,
scalings=cls.scalings,
use_cuda_graph=cls.use_cuda_graph,
)
def test_run_lora_a_sgemm(self):
batch_size = 2
input_dim = 4
output_dim = 6
num_loras = 3
dtype = torch.float32
x = torch.randn(batch_size, input_dim, dtype=dtype)
weights = torch.randn(num_loras, output_dim, input_dim, dtype=dtype)
total_seq_len, _ = x.shape
_, weight_output_dim, _ = weights.shape
output_tensor = torch.zeros(
(total_seq_len, weight_output_dim), dtype=dtype, device=self.device
)
sgmv_shrink(
x,
weights,
output_tensor,
self.backend.batch_info.seg_lens,
self.backend.batch_info.weight_indices,
1.0,
)
scaling = torch.repeat_interleave(
self.backend.batch_info.scalings[self.backend.batch_info.weight_indices],
self.backend.batch_info.seg_lens,
output_size=total_seq_len,
).unsqueeze(-1)
expect_output = output_tensor * scaling
actual_output = self.backend.run_lora_a_sgemm(x, weights)
self.assertTrue(torch.allclose(actual_output, expect_output))
def test_run_lora_b_sgemm(self):
batch_size = 2
input_dim = 6
output_dim = 4
num_loras = 3
dtype = torch.float32
x = torch.randn(batch_size, input_dim, dtype=dtype)
weights = torch.randn(num_loras, output_dim, input_dim, dtype=dtype)
total_seq_len, _ = x.shape
_, weight_output_dim, _ = weights.shape
output_tensor = torch.zeros(
(total_seq_len, weight_output_dim), dtype=dtype, device=self.device
)
sgmv_expand(
x,
weights,
output_tensor,
self.backend.batch_info.seg_lens,
self.backend.batch_info.weight_indices,
True,
)
expect_output = output_tensor
actual_output = self.backend.run_lora_b_sgemm(x, weights)
self.assertTrue(torch.allclose(actual_output, expect_output))
def test_run_qkv_lora(self):
batch_size = 2
input_dim = 6
output_dim = 4
num_loras = 3
dtype = torch.float32
x = torch.randn(batch_size, input_dim, dtype=dtype)
qkv_lora_a = torch.randn(num_loras, output_dim, input_dim, dtype=dtype)
qkv_lora_b = torch.randn(num_loras, input_dim, output_dim, dtype=dtype)
output_offset_cpu = torch.tensor([0, 3, 6, 9, 12], dtype=torch.int32)
num_slices = 3
total_seq_len, _ = x.shape
_, weight_intermediate_dim, _ = qkv_lora_a.shape
_, weight_out_dim, _ = qkv_lora_b.shape
max_rank = weight_intermediate_dim // num_slices
output_tensor = torch.zeros(
(total_seq_len, weight_out_dim), device=x.device, dtype=x.dtype
)
lora_a_output = torch.zeros(
total_seq_len, weight_intermediate_dim, dtype=x.dtype, device=x.device
)
sgmv_shrink(
x,
qkv_lora_a,
lora_a_output,
self.backend.batch_info.seg_lens,
self.backend.batch_info.weight_indices,
1.0,
)
scaling = torch.repeat_interleave(
self.backend.batch_info.scalings[self.backend.batch_info.weight_indices],
self.backend.batch_info.seg_lens,
output_size=total_seq_len,
).unsqueeze(-1)
lora_a_output = lora_a_output * scaling
for slice_id in range(num_slices):
slice_offset = output_offset_cpu[slice_id]
slice_offset_next = output_offset_cpu[slice_id + 1]
slice_size = slice_offset_next - slice_offset
sgmv_expand_slice(
lora_a_output[:, (max_rank * slice_id) : (max_rank * (slice_id + 1))],
qkv_lora_b[:, slice_offset:slice_offset_next],
output_tensor,
self.backend.batch_info.seg_lens,
self.backend.batch_info.weight_indices,
slice_offset,
slice_size,
True,
)
expect_output = output_tensor
actual_output = self.backend.run_qkv_lora(
x, qkv_lora_a, qkv_lora_b, None, output_offset_cpu, 0
)
self.assertTrue(torch.allclose(actual_output, expect_output))
def test_run_gate_up_lora(self):
batch_size = 2
input_dim = 6
output_dim = 4
num_loras = 3
dtype = torch.float32
num_slices = 2
x = torch.randn(batch_size, input_dim, dtype=dtype)
gate_up_lora_a = torch.randn(num_loras, output_dim, input_dim, dtype=dtype)
gate_up_lora_b = torch.randn(
num_loras, output_dim, output_dim // num_slices, dtype=dtype
)
total_seq_len, _ = x.shape
_, weight_intermediate_dim, _ = gate_up_lora_a.shape
_, weight_out_dim, _ = gate_up_lora_b.shape
slice_size = weight_out_dim // num_slices
max_rank = weight_intermediate_dim // num_slices
output_tensor = torch.zeros(
(total_seq_len, weight_out_dim), device=x.device, dtype=x.dtype
)
lora_a_output = torch.zeros(
total_seq_len, weight_intermediate_dim, dtype=x.dtype, device=x.device
)
sgmv_shrink(
x,
gate_up_lora_a,
lora_a_output,
self.backend.batch_info.seg_lens,
self.backend.batch_info.weight_indices,
1.0,
)
scaling = torch.repeat_interleave(
self.backend.batch_info.scalings[self.backend.batch_info.weight_indices],
self.backend.batch_info.seg_lens,
output_size=total_seq_len,
).unsqueeze(-1)
lora_a_output = lora_a_output * scaling
slice_offset = 0
for slice_id in range(num_slices):
sgmv_expand_slice(
lora_a_output[:, (max_rank * slice_id) : (max_rank * (slice_id + 1))],
gate_up_lora_b[:, slice_offset : slice_offset + slice_size],
output_tensor,
self.backend.batch_info.seg_lens,
self.backend.batch_info.weight_indices,
slice_offset,
slice_size,
True,
)
slice_offset += slice_size
expect_output = output_tensor
actual_output = self.backend.run_gate_up_lora(x, gate_up_lora_a, gate_up_lora_b)
self.assertTrue(torch.allclose(actual_output, expect_output))
if __name__ == "__main__":
unittest.main()