From be0cca5596f97c2eff6f5e7b4ba49f187e7a2e11 Mon Sep 17 00:00:00 2001 From: satyamk7054 <43010011+satyamk7054@users.noreply.github.com> Date: Thu, 26 Mar 2026 14:35:20 -0700 Subject: [PATCH] Use torch.addmm instead of separate mm and add_ calls for LoRA torch.native (#20562) Co-authored-by: Satyam Kumar --- .../sglang/srt/lora/backend/torch_backend.py | 10 +++++++--- python/sglang/srt/lora/torch_ops/lora_ops.py | 19 +++++++++++++------ 2 files changed, 20 insertions(+), 9 deletions(-) diff --git a/python/sglang/srt/lora/backend/torch_backend.py b/python/sglang/srt/lora/backend/torch_backend.py index 51b19fa4b..64337acce 100644 --- a/python/sglang/srt/lora/backend/torch_backend.py +++ b/python/sglang/srt/lora/backend/torch_backend.py @@ -23,6 +23,9 @@ class TorchNativeLoRABatchInfo(LoRABatchInfo): # The index of lora adapter used by each segment, in shape (num_segments,) placed on cpu device weight_indices_cpu: Optional[torch.Tensor] = None + # Scaling factors for each lora adapter, in shape (lora_num,) placed on cpu device + scalings_cpu: Optional[torch.Tensor] = None + class TorchNativeLoRABackend(BaseLoRABackend): name = "torch_native" @@ -44,7 +47,7 @@ class TorchNativeLoRABackend(BaseLoRABackend): 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, + scaling_tensor=self.batch_info.scalings_cpu, num_slices=1, ) @@ -93,7 +96,7 @@ class TorchNativeLoRABackend(BaseLoRABackend): 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, + scaling_tensor=self.batch_info.scalings_cpu, num_slices=num_slices, ) @@ -131,7 +134,7 @@ class TorchNativeLoRABackend(BaseLoRABackend): 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, + scaling_tensor=self.batch_info.scalings_cpu, num_slices=num_slices, ) @@ -274,5 +277,6 @@ class TorchNativeLoRABackend(BaseLoRABackend): batch_info.seg_indptr_cpu = seg_indptr_cpu batch_info.seg_lens_cpu = seg_lens_cpu batch_info.weight_indices_cpu = weight_indices_tensor + batch_info.scalings_cpu = scalings_tensor self.batch_info = batch_info diff --git a/python/sglang/srt/lora/torch_ops/lora_ops.py b/python/sglang/srt/lora/torch_ops/lora_ops.py index 562956327..235ed0513 100644 --- a/python/sglang/srt/lora/torch_ops/lora_ops.py +++ b/python/sglang/srt/lora/torch_ops/lora_ops.py @@ -31,13 +31,19 @@ def sgemm_lora_a_fwd( continue if rank > 0: - x_seq = inputs[token_offset : token_offset + seq_len, :] w_seq = weights[lora_idx, : num_slices * rank, :] - result = torch.mm(x_seq, w_seq.T) - output[token_offset : token_offset + seq_len, : num_slices * rank] = ( - scaling_tensor[lora_idx] * result + out_slice = output[ + token_offset : token_offset + seq_len, : num_slices * rank + ] + torch.addmm( + out_slice, + x_seq, + w_seq.T, + beta=0, + alpha=scaling_tensor[lora_idx].item(), + out=out_slice, ) token_offset += seq_len @@ -98,10 +104,11 @@ def sgemm_lora_b_fwd( lora_idx, slice_start_output:slice_end_output, :rank ] # (slice_dim, rank) - output[ + out_slice = output[ token_offset : token_offset + seq_len, slice_start_output:slice_end_output, - ].add_(torch.mm(x_slice, w_slice.T)) + ] + torch.addmm(out_slice, x_slice, w_slice.T, beta=1, alpha=1, out=out_slice) token_offset += seq_len