diff --git a/python/sglang/srt/lora/backend/torch_backend.py b/python/sglang/srt/lora/backend/torch_backend.py index 3605d29e9..51b19fa4b 100644 --- a/python/sglang/srt/lora/backend/torch_backend.py +++ b/python/sglang/srt/lora/backend/torch_backend.py @@ -222,6 +222,7 @@ class TorchNativeLoRABackend(BaseLoRABackend): ) bs = forward_batch.batch_size + num_segments = len(weight_indices_tensor) if use_cuda_graph: assert ( @@ -229,13 +230,13 @@ class TorchNativeLoRABackend(BaseLoRABackend): ), "CUDA Graph batch info is not initialized." batch_info = self.cuda_graph_batch_info batch_info.bs = forward_batch.batch_size - batch_info.num_segments = forward_batch.batch_size + batch_info.num_segments = num_segments else: max_len = max(seg_lens_cpu) batch_info = TorchNativeLoRABatchInfo( bs=forward_batch.batch_size, - num_segments=forward_batch.batch_size, + num_segments=num_segments, max_len=max_len, use_cuda_graph=False, seg_lens=torch.empty((bs,), dtype=torch.int32, device=self.device), @@ -261,7 +262,9 @@ class TorchNativeLoRABackend(BaseLoRABackend): batch_info.scalings[: self.max_loras_per_batch].copy_( scalings_tensor, non_blocking=True ) - batch_info.weight_indices[:bs].copy_(weight_indices_tensor, non_blocking=True) + batch_info.weight_indices[:num_segments].copy_( + weight_indices_tensor, non_blocking=True + ) batch_info.seg_indptr[: len(seg_indptr_cpu)].copy_( seg_indptr_cpu, non_blocking=True ) diff --git a/test/manual/lora/test_torch_backend.py b/test/manual/lora/test_torch_backend.py index 0e9996de0..f0c46dd3e 100644 --- a/test/manual/lora/test_torch_backend.py +++ b/test/manual/lora/test_torch_backend.py @@ -11,20 +11,22 @@ from sglang.test.test_utils import CustomTestCase class TestTorchNativeLoRABackend(CustomTestCase): device = "cpu" - weight_indices = [0, 1] + + # set duplicate weights to test merging during prepare_lora_batch + weight_indices = [0, 0, 1] lora_ranks = [1, 1] scalings = [1.0, 0.5] - seq_lens = [1, 1] + seq_lens = [1, 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), + batch_size=3, + input_ids=torch.tensor([[1], [2], [3]], dtype=torch.int32), req_pool_indices=None, seq_lens=None, out_cache_loc=None, - seq_lens_sum=6, + seq_lens_sum=3, extend_seq_lens=torch.tensor(seq_lens, dtype=torch.int32), extend_seq_lens_cpu=seq_lens, ) @@ -41,7 +43,7 @@ class TestTorchNativeLoRABackend(CustomTestCase): ) def test_run_lora_a_sgemm(self): - batch_size = 2 + batch_size = 3 input_dim = 4 output_dim = 6 num_loras = 3 @@ -80,7 +82,7 @@ class TestTorchNativeLoRABackend(CustomTestCase): self.assertTrue(torch.allclose(actual_output, expect_output)) def test_run_lora_b_sgemm(self): - batch_size = 2 + batch_size = 3 input_dim = 6 output_dim = 4 num_loras = 3 @@ -118,12 +120,12 @@ class TestTorchNativeLoRABackend(CustomTestCase): self.assertTrue(torch.allclose(actual_output, expect_output)) def test_run_qkv_lora(self): - batch_size = 2 + batch_size = 3 num_loras = 3 input_dim = 6 - output_offset = [0, 3, 6, 9, 12] + output_offset = [0, 3, 6, 9] output_dim = output_offset[-1] - num_slices = len(output_offset) - 1 + num_slices = len(output_offset) - 1 # 3 slices for Q, K, V max_lora_rank = max(self.lora_ranks) dtype = torch.float32 @@ -177,7 +179,7 @@ class TestTorchNativeLoRABackend(CustomTestCase): self.assertTrue(torch.allclose(actual_output, expect_output)) def test_run_gate_up_lora(self): - batch_size = 2 + batch_size = 3 input_dim = 6 output_dim = 4 num_loras = 3