From 263cb3b222f6e1ce106600bbf783cb85a9e577b6 Mon Sep 17 00:00:00 2001 From: Vladimir Serov Date: Thu, 7 May 2026 17:28:38 +0300 Subject: [PATCH] [LoRA] Torch Native enhancement: embedding and graph optimization (#21885) Co-authored-by: ronnie_zheng --- .../sglang/srt/lora/backend/torch_backend.py | 118 ++++--- python/sglang/srt/lora/layers.py | 88 ++++-- python/sglang/srt/lora/torch_ops/__init__.py | 105 ++++++- .../srt/lora/torch_ops/graph_lora_ops.py | 120 +++++++ python/sglang/srt/lora/torch_ops/lora_ops.py | 103 +++--- python/sglang/test/lora_utils.py | 7 +- test/manual/lora/test_lora_ops.py | 296 +++++++++++++++++- .../lora/test_chunked_sgmv_backend.py | 1 + 8 files changed, 716 insertions(+), 122 deletions(-) create mode 100644 python/sglang/srt/lora/torch_ops/graph_lora_ops.py diff --git a/python/sglang/srt/lora/backend/torch_backend.py b/python/sglang/srt/lora/backend/torch_backend.py index 13f0b9ba0..f9f938e87 100644 --- a/python/sglang/srt/lora/backend/torch_backend.py +++ b/python/sglang/srt/lora/backend/torch_backend.py @@ -4,7 +4,11 @@ from typing import Optional import torch from sglang.srt.lora.backend.base_backend import BaseLoRABackend -from sglang.srt.lora.torch_ops import sgemm_lora_a_fwd, sgemm_lora_b_fwd +from sglang.srt.lora.torch_ops import ( + sgemm_lora_a_embedding_fwd, + sgemm_lora_a_fwd, + sgemm_lora_b_fwd, +) from sglang.srt.lora.utils import LoRABatchInfo, generate_sequence_lengths from sglang.srt.model_executor.forward_batch_info import ForwardBatch @@ -38,6 +42,27 @@ class TorchNativeLoRABackend(BaseLoRABackend): ): super().__init__(max_loras_per_batch, device) + def run_lora_a_embedding( + self, + input_ids: torch.Tensor, + weights: torch.Tensor, + vocab_size: int, + extra_embeddings: torch.Tensor = None, + *args, + **kwargs, + ) -> torch.Tensor: + assert ( + extra_embeddings is None + ), "Extra embeddings for lora a is not supported yet in chunked backend" + output_tensor = sgemm_lora_a_embedding_fwd( + inputs=input_ids, + weights=weights, + batch_info=self.batch_info, + vocab_size=vocab_size, + ) + + return output_tensor + def run_lora_a_sgemm( self, x: torch.Tensor, @@ -49,10 +74,7 @@ class TorchNativeLoRABackend(BaseLoRABackend): output_tensor = sgemm_lora_a_fwd( inputs=x, weights=weights, - 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, - scaling_tensor=self.batch_info.scalings_cpu, + batch_info=self.batch_info, num_slices=stack_num, ) @@ -62,21 +84,18 @@ class TorchNativeLoRABackend(BaseLoRABackend): self, x: torch.Tensor, weights: torch.Tensor, + output_offset_cpu: torch.Tensor, base_output: torch.Tensor = None, *args, **kwargs, ) -> torch.Tensor: _, weight_out_dim, _ = weights.shape - output_offset = torch.tensor( - [0, weight_out_dim], dtype=torch.int32, device="cpu" - ) + output_tensor = sgemm_lora_b_fwd( inputs=x, weights=weights, - 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, + batch_info=self.batch_info, + slice_offsets=output_offset_cpu, base_output=base_output, ) @@ -98,19 +117,14 @@ class TorchNativeLoRABackend(BaseLoRABackend): lora_a_output = sgemm_lora_a_fwd( inputs=x, weights=qkv_lora_a, - 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, - scaling_tensor=self.batch_info.scalings_cpu, + batch_info=self.batch_info, num_slices=n_slices, ) output_tensor = sgemm_lora_b_fwd( inputs=lora_a_output, weights=qkv_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, + batch_info=self.batch_info, slice_offsets=output_offset_cpu, base_output=base_output, ) @@ -122,34 +136,26 @@ class TorchNativeLoRABackend(BaseLoRABackend): x: torch.Tensor, gate_up_lora_a: torch.Tensor, gate_up_lora_b: torch.Tensor, + output_offset_cpu: torch.Tensor, base_output: torch.Tensor = None, *args, **kwargs, ) -> torch.Tensor: - num_slices = 2 + num_slices = len(output_offset_cpu) - 1 _, 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" - ) lora_a_output = sgemm_lora_a_fwd( inputs=x, weights=gate_up_lora_a, - 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, - scaling_tensor=self.batch_info.scalings_cpu, + batch_info=self.batch_info, num_slices=num_slices, ) output_tensor = sgemm_lora_b_fwd( inputs=lora_a_output, 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, + batch_info=self.batch_info, + slice_offsets=output_offset_cpu, base_output=base_output, ) @@ -192,36 +198,44 @@ class TorchNativeLoRABackend(BaseLoRABackend): scalings: list[float], use_cuda_graph: bool, ): + # Do not use merge optimization for graph mode + # Use pinned memory to avoid synchronizations during host-to-device transfer original_seq_lens_cpu = generate_sequence_lengths(forward_batch, device="cpu") - original_weight_indices_tensor = torch.tensor( - 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 + if not use_cuda_graph: + original_weight_indices_tensor = torch.tensor( + weight_indices, dtype=torch.int32, device="cpu" ) - ) - seg_lens_cpu = ( - torch.zeros_like( - unique_weight_indices_tensor, dtype=torch.int32, device="cpu" + unique_weight_indices_tensor, inverse_weight_indices_tensor = ( + torch.unique_consecutive( + original_weight_indices_tensor, return_inverse=True + ) ) - .scatter_add_( - 0, - inverse_weight_indices_tensor, + + 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() + ) + + weight_indices_tensor = unique_weight_indices_tensor.pin_memory() + else: + weight_indices_tensor = torch.repeat_interleave( + torch.tensor(weight_indices, dtype=torch.int32, device="cpu"), original_seq_lens_cpu, - ) - .pin_memory() - ) + ).pin_memory() + seg_lens_cpu = torch.ones_like(weight_indices_tensor).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, dtype=torch.int32, pin_memory=True, device="cpu" ) diff --git a/python/sglang/srt/lora/layers.py b/python/sglang/srt/lora/layers.py index 41eeae3ee..475df0067 100644 --- a/python/sglang/srt/lora/layers.py +++ b/python/sglang/srt/lora/layers.py @@ -90,12 +90,18 @@ class VocabParallelEmbeddingWithLoRA(BaseLayerWithLoRA): assert ( not get_attn_tp_context().allow_input_scattered ), "VocabParallelEmbeddingWithLoRA with TP > 1 under input_scattered mode (e.g., DeepSeek-v2 MLA with --enable-attn-tp-input-scattered) is not fully supported and may produce incorrect results. Consider disabling input_scattered or removing embed_tokens from LoRA target modules." - + offsets = [0, self.embed_dim] self.output_offset = torch.tensor( - [0, self.embed_dim], + offsets, dtype=torch.int32, device=next(base_layer.parameters()).device, ) + self.output_offset_cpu = torch.tensor( + offsets, + dtype=torch.int32, + device="cpu", + pin_memory=True, + ) def set_lora_info( self, @@ -125,6 +131,7 @@ class VocabParallelEmbeddingWithLoRA(BaseLayerWithLoRA): x=lora_a_output, weights=self.embedding_B_buffer, output_offset=self.output_offset, + output_offset_cpu=self.output_offset_cpu, base_output=base_output, ) return lora_output @@ -243,6 +250,8 @@ class ParallelLMHeadWithLoRA(BaseLayerWithLoRA): self.embed_dim = base_layer.embedding_dim self.vocab_size = base_layer.org_vocab_size + offsets = [0, self.vocab_size] + tp_size = base_layer.tp_size if hasattr(base_layer, "tp_size") else 1 # lm_head LoRA keeps A unsharded and shards B along the vocab @@ -265,17 +274,19 @@ class ParallelLMHeadWithLoRA(BaseLayerWithLoRA): self.vocab_size, shard_indices=base_layer.shard_indices, ) - self.output_offset = torch.tensor( - [0, self.shard_vocab_size], - dtype=torch.int32, - device=next(base_layer.parameters()).device, - ) - else: - self.output_offset = torch.tensor( - [0, self.vocab_size], - dtype=torch.int32, - device=next(base_layer.parameters()).device, - ) + offsets = [0, self.shard_vocab_size] + + self.output_offset = torch.tensor( + offsets, + dtype=torch.int32, + device=next(base_layer.parameters()).device, + ) + self.output_offset_cpu = torch.tensor( + offsets, + dtype=torch.int32, + device="cpu", + pin_memory=True, + ) def set_lora_info( self, @@ -349,6 +360,7 @@ class ParallelLMHeadWithLoRA(BaseLayerWithLoRA): x=lora_a_output, weights=self.lm_head_B_buffer, output_offset=self.output_offset, + output_offset_cpu=self.output_offset_cpu, base_output=base_output, pruned_batch_info=lm_head_batch_info, ) @@ -411,14 +423,18 @@ class ColumnParallelLinearWithLoRA(BaseLayerWithLoRA): ) -> None: super().__init__(base_layer, lora_backend) shard_size = self.base_layer.output_partition_sizes[0] + offsets = [0, shard_size] self.output_offset = torch.tensor( - [ - 0, - shard_size, - ], + offsets, dtype=torch.int32, device=next(self.base_layer.parameters()).device, ) + self.output_offset_cpu = torch.tensor( + offsets, + dtype=torch.int32, + device="cpu", + pin_memory=True, + ) def set_lora_info( self, @@ -435,6 +451,7 @@ class ColumnParallelLinearWithLoRA(BaseLayerWithLoRA): x=lora_a_output, weights=self.B_buffer, output_offset=self.output_offset, + output_offset_cpu=self.output_offset_cpu, base_output=base_output, ) return lora_output @@ -510,7 +527,7 @@ class MergedColumnParallelLinearWithLoRA(ColumnParallelLinearWithLoRA): dtype=torch.int32, device=next(self.base_layer.parameters()).device, ) - self.output_offset_cpu = self.output_offset.cpu() + self.output_offset_cpu = self.output_offset.cpu().pin_memory() self.max_out_dim = max(partition_sizes) self.use_gate_up_lora = ( lora_n_slices == 2 and partition_sizes[0] == partition_sizes[1] @@ -536,6 +553,7 @@ class MergedColumnParallelLinearWithLoRA(ColumnParallelLinearWithLoRA): gate_up_lora_a=self.A_buffer, gate_up_lora_b=self.B_buffer, output_offset=self.output_offset, + output_offset_cpu=self.output_offset_cpu, base_output=base_output, ) else: @@ -576,17 +594,23 @@ class QKVParallelLinearWithLoRA(ColumnParallelLinearWithLoRA): super().__init__(base_layer, lora_backend) q_proj_shard_size = self.base_layer.q_proj_shard_size kv_proj_shard_size = self.base_layer.kv_proj_shard_size + offsets = [ + 0, + q_proj_shard_size, + q_proj_shard_size + kv_proj_shard_size, + q_proj_shard_size + 2 * kv_proj_shard_size, + ] self.output_offset = torch.tensor( - [ - 0, - q_proj_shard_size, - q_proj_shard_size + kv_proj_shard_size, - q_proj_shard_size + 2 * kv_proj_shard_size, - ], + offsets, dtype=torch.int32, device=next(self.base_layer.parameters()).device, ) - self.output_offset_cpu = self.output_offset.cpu() + self.output_offset_cpu = torch.tensor( + offsets, + dtype=torch.int32, + device="cpu", + pin_memory=True, + ) # For computing number of launched blocks self.max_qkv_out_dim = max(q_proj_shard_size, kv_proj_shard_size) @@ -658,14 +682,18 @@ class RowParallelLinearWithLoRA(BaseLayerWithLoRA): self.A_buffer = A_buffer self.B_buffer = B_buffer output_size = self.base_layer.output_size + offsets = [0, output_size] self.output_offset = torch.tensor( - [ - 0, - output_size, - ], + offsets, dtype=torch.int32, device=next(self.base_layer.parameters()).device, ) + self.output_offset_cpu = torch.tensor( + offsets, + dtype=torch.int32, + device="cpu", + pin_memory=True, + ) def apply_lora(self, base_output: torch.Tensor, x: torch.Tensor) -> torch.Tensor: lora_a_output = self.lora_backend.run_lora_a_sgemm(x, self.A_buffer) @@ -673,6 +701,7 @@ class RowParallelLinearWithLoRA(BaseLayerWithLoRA): x=lora_a_output, weights=self.B_buffer, output_offset=self.output_offset, + output_offset_cpu=self.output_offset_cpu, base_output=base_output, ) return lora_output @@ -712,6 +741,7 @@ class RowParallelLinearWithLoRA(BaseLayerWithLoRA): x=lora_a_output, weights=self.B_buffer, output_offset=self.output_offset, + output_offset_cpu=self.output_offset_cpu, base_output=output_, ) else: diff --git a/python/sglang/srt/lora/torch_ops/__init__.py b/python/sglang/srt/lora/torch_ops/__init__.py index bc3a5391d..807c1d0a9 100644 --- a/python/sglang/srt/lora/torch_ops/__init__.py +++ b/python/sglang/srt/lora/torch_ops/__init__.py @@ -1,6 +1,109 @@ -from .lora_ops import sgemm_lora_a_fwd, sgemm_lora_b_fwd +from typing import Optional + +import torch + +from sglang.srt.lora.utils import LoRABatchInfo + +from .graph_lora_ops import ( + sgemm_lora_a_embedding_graph_fwd, + sgemm_lora_a_graph_fwd, + sgemm_lora_b_graph_fwd, +) +from .lora_ops import sgemm_lora_a_embedding_fwd as sgemm_lora_a_embedding_control_fwd +from .lora_ops import sgemm_lora_a_fwd as sgemm_lora_a_control_fwd +from .lora_ops import sgemm_lora_b_fwd as sgemm_lora_b_control_fwd + + +def sgemm_lora_a_embedding_fwd( + inputs: torch.Tensor, + weights: torch.Tensor, + batch_info: LoRABatchInfo, + vocab_size: int, +) -> torch.Tensor: + output: torch.Tensor + if batch_info.use_cuda_graph: + output = sgemm_lora_a_embedding_graph_fwd( + inputs, + weights, + batch_info.weight_indices, + batch_info.seg_lens, + batch_info.scalings, + vocab_size, + ) + else: + output = sgemm_lora_a_embedding_control_fwd( + inputs, + weights, + batch_info.weight_indices_cpu, + batch_info.seg_lens_cpu, + batch_info.lora_ranks_cpu, + batch_info.scalings_cpu, + vocab_size, + ) + return output + + +def sgemm_lora_a_fwd( + inputs: torch.Tensor, + weights: torch.Tensor, + batch_info: LoRABatchInfo, + num_slices: int = 1, +) -> torch.Tensor: + output: torch.Tensor + if batch_info.use_cuda_graph: + output = sgemm_lora_a_graph_fwd( + inputs, + weights, + batch_info.weight_indices, + batch_info.seg_lens, + batch_info.scalings, + num_slices, + ) + else: + output = sgemm_lora_a_control_fwd( + inputs, + weights, + batch_info.weight_indices_cpu, + batch_info.seg_lens_cpu, + batch_info.lora_ranks_cpu, + batch_info.scalings_cpu, + num_slices, + ) + return output + + +def sgemm_lora_b_fwd( + inputs: torch.Tensor, + weights: torch.Tensor, + batch_info: LoRABatchInfo, + slice_offsets: torch.Tensor, + base_output: Optional[torch.Tensor] = None, +) -> torch.Tensor: + output: torch.Tensor + if batch_info.use_cuda_graph: + output = sgemm_lora_b_graph_fwd( + inputs, + weights, + batch_info.weight_indices, + batch_info.seg_lens, + slice_offsets, + base_output, + ) + else: + output = sgemm_lora_b_control_fwd( + inputs, + weights, + batch_info.weight_indices_cpu, + batch_info.seg_lens_cpu, + batch_info.lora_ranks_cpu, + slice_offsets, + base_output, + ) + return output + __all__ = [ + "sgemm_lora_a_embedding_fwd", "sgemm_lora_a_fwd", "sgemm_lora_b_fwd", ] diff --git a/python/sglang/srt/lora/torch_ops/graph_lora_ops.py b/python/sglang/srt/lora/torch_ops/graph_lora_ops.py new file mode 100644 index 000000000..4317bc894 --- /dev/null +++ b/python/sglang/srt/lora/torch_ops/graph_lora_ops.py @@ -0,0 +1,120 @@ +from typing import Optional + +import torch +import torch.nn.functional as F + + +def sgemm_lora_a_embedding_graph_fwd( + inputs: torch.Tensor, + weights: torch.Tensor, + weight_indices: torch.Tensor, + seg_len_tensor: torch.Tensor, + scaling_tensor: torch.Tensor, + vocab_size: int, +) -> torch.Tensor: + total_seq_len = inputs.shape[0] + if weights.numel() == 0: + return torch.zeros(total_seq_len, 0, dtype=weights.dtype, device=weights.device) + + num_loras, max_rank, _ = weights.shape + + output = torch.zeros( + total_seq_len, max_rank, dtype=weights.dtype, device=weights.device + ) + + for lora_idx in range(num_loras): + + batch_token_mask = weight_indices[:total_seq_len] == lora_idx + + x_seq = torch.where(batch_token_mask, inputs, 0) + w_seq = weights[lora_idx] + + output.add_( + scaling_tensor[lora_idx] + * torch.where( + batch_token_mask.unsqueeze(1), F.embedding(x_seq, w_seq.t()), 0 + ) + ) + + return output + + +def sgemm_lora_a_graph_fwd( + inputs: torch.Tensor, + weights: torch.Tensor, + weight_indices: torch.Tensor, + seg_len_tensor: torch.Tensor, + scaling_tensor: torch.Tensor, + num_slices: int = 1, +) -> torch.Tensor: + total_seq_len, input_dim = inputs.shape + if weights.numel() == 0: + 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 + ) + + for lora_idx in range(num_loras): + + batch_token_mask = (weight_indices[:total_seq_len] == lora_idx).unsqueeze(1) + + x_seq = torch.where(batch_token_mask, inputs, 0) + w_seq = weights[lora_idx] + + output.add_(scaling_tensor[lora_idx] * torch.mm(x_seq, w_seq.t(), 0)) + + return output + + +def sgemm_lora_b_graph_fwd( + inputs: torch.Tensor, + weights: torch.Tensor, + weight_indices: torch.Tensor, + seg_len_tensor: torch.Tensor, + slice_offsets: torch.Tensor, + base_output: Optional[torch.Tensor] = None, +) -> torch.Tensor: + total_seq_len, input_dim = inputs.shape + num_loras, weight_out_dim, _ = weights.shape + total_output_dim = slice_offsets[-1].item() if len(slice_offsets) > 0 else 0 + + if weights.numel() == 0: + return torch.zeros( + total_seq_len, total_output_dim, dtype=inputs.dtype, device=inputs.device + ) + + num_slices = len(slice_offsets) - 1 + max_rank = input_dim // num_slices + + if base_output is not None: + output = base_output + else: + output = torch.zeros( + total_seq_len, total_output_dim, dtype=inputs.dtype, device=inputs.device + ) + + for lora_idx in range(num_loras): + + batch_token_mask = (weight_indices[:total_seq_len] == lora_idx).unsqueeze(1) + inputs_masked = torch.where(batch_token_mask, inputs, 0) + + for slice_idx in range(num_slices): + slice_start_input = slice_idx * max_rank + slice_end_input = (slice_idx + 1) * max_rank + + slice_start_output = slice_offsets[slice_idx] + slice_end_output = slice_offsets[slice_idx + 1] + + x_slice = inputs_masked[..., slice_start_input:slice_end_input] + w_slice = weights[ + lora_idx, slice_start_output:slice_end_output + ] # (slice_dim, max_rank) + output[..., slice_start_output:slice_end_output].add_( + torch.mm(x_slice, w_slice.t()) + ) + + return output diff --git a/python/sglang/srt/lora/torch_ops/lora_ops.py b/python/sglang/srt/lora/torch_ops/lora_ops.py index 235ed0513..16b25dc35 100644 --- a/python/sglang/srt/lora/torch_ops/lora_ops.py +++ b/python/sglang/srt/lora/torch_ops/lora_ops.py @@ -3,6 +3,46 @@ from typing import Optional import torch +def sgemm_lora_a_embedding_fwd( + inputs: torch.Tensor, + weights: torch.Tensor, + weight_indices: torch.Tensor, + seg_len_tensor: torch.Tensor, + lora_ranks: torch.Tensor, + scaling_tensor: torch.Tensor, + vocab_size: int, +) -> torch.Tensor: + total_seq_len = inputs.shape[0] + if weights.numel() == 0: + return torch.zeros(total_seq_len, 0, dtype=weights.dtype, device=weights.device) + + num_loras, max_rank, _ = weights.shape + + output = torch.zeros( + total_seq_len, max_rank, dtype=weights.dtype, device=weights.device + ) + + token_offset = 0 + for lora_idx, seq_len in zip(weight_indices, seg_len_tensor): + if seq_len == 0: + continue + + rank = lora_ranks[lora_idx] + if rank > 0: + + x_seq = inputs[token_offset : token_offset + seq_len] + w_seq = weights[lora_idx, :rank] + + result = torch.nn.functional.embedding(x_seq, w_seq.T) + output[token_offset : token_offset + seq_len, :rank] = ( + scaling_tensor[lora_idx].item() * result + ) + + token_offset += seq_len + + return output + + def sgemm_lora_a_fwd( inputs: torch.Tensor, weights: torch.Tensor, @@ -11,7 +51,7 @@ def sgemm_lora_a_fwd( lora_ranks: torch.Tensor, scaling_tensor: torch.Tensor, num_slices: int = 1, -): +) -> torch.Tensor: total_seq_len, input_dim = inputs.shape if weights.numel() == 0: return torch.zeros(total_seq_len, 0, dtype=inputs.dtype, device=inputs.device) @@ -24,26 +64,21 @@ def sgemm_lora_a_fwd( ) token_offset = 0 - for lora_idx, seq_len, rank in zip( - weight_indices, seg_len_tensor, lora_ranks[weight_indices] - ): + for lora_idx, seq_len in zip(weight_indices, seg_len_tensor): if seq_len == 0: continue + rank = lora_ranks[lora_idx] if rank > 0: - x_seq = inputs[token_offset : token_offset + seq_len, :] - w_seq = weights[lora_idx, : num_slices * rank, :] - out_slice = output[ - token_offset : token_offset + seq_len, : num_slices * rank - ] - torch.addmm( - out_slice, + x_seq = inputs[token_offset : token_offset + seq_len] + w_seq = weights[lora_idx, : num_slices * rank] + + output[token_offset : token_offset + seq_len, : num_slices * rank].addmm_( x_seq, w_seq.T, beta=0, alpha=scaling_tensor[lora_idx].item(), - out=out_slice, ) token_offset += seq_len @@ -59,7 +94,7 @@ def sgemm_lora_b_fwd( lora_ranks: torch.Tensor, slice_offsets: torch.Tensor, base_output: Optional[torch.Tensor] = None, -): +) -> torch.Tensor: total_seq_len, _ = inputs.shape num_loras, weight_out_dim, _ = weights.shape total_output_dim = slice_offsets[-1].item() if len(slice_offsets) > 0 else 0 @@ -79,36 +114,32 @@ def sgemm_lora_b_fwd( ) token_offset = 0 - for lora_idx, seq_len, rank in zip( - weight_indices, seg_len_tensor, lora_ranks[weight_indices] - ): + for lora_idx, seq_len in zip(weight_indices, seg_len_tensor): if seq_len == 0: continue - if rank == 0: - token_offset += seq_len - continue + rank = lora_ranks[lora_idx] + if rank > 0: - for slice_idx in range(num_slices): - slice_start_input = slice_idx * rank - slice_end_input = (slice_idx + 1) * 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] - slice_end_output = slice_offsets[slice_idx + 1] + slice_start_output = slice_offsets[slice_idx] + slice_end_output = slice_offsets[slice_idx + 1] - x_slice = inputs[ - token_offset : token_offset + seq_len :, - slice_start_input:slice_end_input, - ] # (seq_len, rank) - w_slice = weights[ - lora_idx, slice_start_output:slice_end_output, :rank - ] # (slice_dim, rank) + x_slice = inputs[ + token_offset : token_offset + seq_len, + slice_start_input:slice_end_input, + ] # (seq_len, rank) + w_slice = weights[ + lora_idx, slice_start_output:slice_end_output, :rank + ] # (slice_dim, rank) - out_slice = output[ - token_offset : token_offset + seq_len, - slice_start_output:slice_end_output, - ] - torch.addmm(out_slice, x_slice, w_slice.T, beta=1, alpha=1, out=out_slice) + output[ + token_offset : token_offset + seq_len, + slice_start_output:slice_end_output, + ].addmm_(x_slice, w_slice.T) token_offset += seq_len diff --git a/python/sglang/test/lora_utils.py b/python/sglang/test/lora_utils.py index 39258ea18..0bfd7fa07 100644 --- a/python/sglang/test/lora_utils.py +++ b/python/sglang/test/lora_utils.py @@ -236,6 +236,7 @@ def reference_embedding_lora_a_shrink( weight_indices: torch.Tensor, seq_lengths: torch.Tensor, lora_ranks: torch.Tensor, + lora_scalings: torch.Tensor, vocab_size: int, ) -> torch.Tensor: """ @@ -247,6 +248,7 @@ def reference_embedding_lora_a_shrink( 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 vocab_size: Base vocabulary size Returns: @@ -264,10 +266,11 @@ def reference_embedding_lora_a_shrink( ) token_offset = 0 - for lora_idx, seq_len, rank in zip( + 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 @@ -284,7 +287,7 @@ def reference_embedding_lora_a_shrink( lora_weights = weights[lora_idx, :rank, :] # (rank, vocab_size) embeddings = lora_weights[:, clamped_ids].t() # (seq_len, rank) - output[token_offset : token_offset + seq_len, :rank] = embeddings + output[token_offset : token_offset + seq_len, :rank] = scaling * embeddings token_offset += seq_len diff --git a/test/manual/lora/test_lora_ops.py b/test/manual/lora/test_lora_ops.py index 5ed4b8ed7..385500056 100644 --- a/test/manual/lora/test_lora_ops.py +++ b/test/manual/lora/test_lora_ops.py @@ -3,12 +3,81 @@ 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.srt.lora.torch_ops.graph_lora_ops import ( + sgemm_lora_a_embedding_graph_fwd, + sgemm_lora_a_graph_fwd, + sgemm_lora_b_graph_fwd, +) +from sglang.srt.lora.torch_ops.lora_ops import ( + sgemm_lora_a_embedding_fwd, + sgemm_lora_a_fwd, + sgemm_lora_b_fwd, +) +from sglang.test.lora_utils import ( + reference_embedding_lora_a_shrink, + reference_sgmv_expand, + reference_sgmv_shrink, +) from sglang.test.test_utils import CustomTestCase class TestLoraOps(CustomTestCase): + def test_sgemm_lora_a_embedding_fwd(self): + batch_size = 64 + input_dim = 1024 + num_loras = 3 + dtype = torch.float32 + vocab_size = 32000 + + 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.randint(vocab_size, (batch_size,), dtype=torch.int32) + lora_a_weights = torch.randn(num_loras, max_lora_rank, vocab_size, 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_embedding_lora_a_shrink( + inputs, + lora_a_weights, + lora_indices_tensor, + seq_len_tensor, + lora_ranks_tensor, + lora_scaling_tensor, + vocab_size, + ) + + actual_output = sgemm_lora_a_embedding_fwd( + inputs, + lora_a_weights, + lora_indices_tensor, + seq_len_tensor, + lora_ranks_tensor, + lora_scaling_tensor, + vocab_size, + ) + + self.assertTrue(torch.allclose(actual_output, expect_output)) + def test_sgemm_lora_a_fwd(self): batch_size = 2 input_dim = 1024 @@ -106,6 +175,67 @@ class TestLoraOps(CustomTestCase): self.assertTrue(torch.allclose(actual_output, expect_output)) + def test_sgemm_lora_a_embedding_fwd_expand(self): + batch_size = 2 + input_dim = 1024 + num_loras = 3 + dtype = torch.float32 + vocab_size = 32000 + + 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.randint(vocab_size, (seq_len,), dtype=torch.int32) + lora_a_weights = torch.randn(num_loras, max_lora_rank, vocab_size, 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_embedding_lora_a_shrink( + inputs, + lora_a_weights, + lora_indices_tensor, + seq_len_tensor, + lora_ranks_tensor, + lora_scaling_tensor, + vocab_size, + ) + + actual_output = sgemm_lora_a_embedding_fwd( + inputs, + lora_a_weights, + lora_indices_tensor, + seq_len_tensor, + lora_ranks_tensor, + lora_scaling_tensor, + vocab_size, + ) + + self.assertTrue(torch.allclose(actual_output, expect_output)) + def test_sgemm_lora_a_fwd_expand(self): batch_size = 2 input_dim = 1024 @@ -213,6 +343,168 @@ class TestLoraOps(CustomTestCase): self.assertTrue(torch.allclose(actual_output, expect_output)) + def test_sgemm_lora_a_embedding_graph_fwd(self): + batch_size = 4 + input_dim = 1024 + num_loras = 3 + dtype = torch.float16 + vocab_size = 32000 + + 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.randint(vocab_size, (batch_size,), dtype=torch.int32) + lora_a_weights = torch.zeros(num_loras, max_lora_rank, vocab_size, dtype=dtype) + for idx, rank in enumerate(lora_ranks): + lora_a_weights[idx, :rank] = torch.randn(rank, vocab_size, 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_embedding_lora_a_shrink( + inputs, + lora_a_weights, + lora_indices_tensor, + seq_len_tensor, + lora_ranks_tensor, + lora_scaling_tensor, + vocab_size, + ) + + actual_output = sgemm_lora_a_embedding_graph_fwd( + inputs, + lora_a_weights, + lora_indices_tensor, + seq_len_tensor, + lora_scaling_tensor, + vocab_size, + ) + + self.assertTrue( + torch.allclose(actual_output, expect_output, rtol=1e-3, atol=1e-5) + ) + + def test_sgemm_lora_a_graph_fwd(self): + batch_size = 4 + input_dim = 1024 + num_loras = 3 + dtype = torch.float16 + + 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.zeros(num_loras, max_lora_rank, input_dim, dtype=dtype) + for idx, rank in enumerate(lora_ranks): + lora_a_weights[idx, :rank] = torch.randn(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_graph_fwd( + inputs, + lora_a_weights, + lora_indices_tensor, + seq_len_tensor, + lora_scaling_tensor, + ) + + self.assertTrue( + torch.allclose(actual_output, expect_output, rtol=1e-3, atol=1e-5) + ) + + def test_sgemm_lora_b_graph_fwd(self): + batch_size = 4 + output_dim = 1024 + num_loras = 3 + dtype = torch.float16 + + 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.zeros(num_loras, output_dim, max_lora_rank, dtype=dtype) + for idx, rank in enumerate(lora_ranks): + lora_b_weights[idx, ..., :rank] = torch.randn(output_dim, 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_graph_fwd( + inputs, + lora_b_weights, + lora_indices_tensor, + seq_len_tensor, + slice_offsets, + ) + + self.assertTrue( + torch.allclose(actual_output, expect_output, rtol=1e-3, atol=1e-5) + ) + if __name__ == "__main__": unittest.main() diff --git a/test/registered/lora/test_chunked_sgmv_backend.py b/test/registered/lora/test_chunked_sgmv_backend.py index fe6c8ed96..253889090 100644 --- a/test/registered/lora/test_chunked_sgmv_backend.py +++ b/test/registered/lora/test_chunked_sgmv_backend.py @@ -528,6 +528,7 @@ class TestChunkedSGMV(unittest.TestCase): lora_assignments_tensor, seq_lengths_tensor, lora_ranks_tensor, + scalings_tensor, self.vocab_size, ) torch.testing.assert_close(