[NPU]LoRA: Adding Torch Native backend (#14132)
This commit is contained in:
287
test/manual/test_lora_ops.py
Normal file
287
test/manual/test_lora_ops.py
Normal file
@@ -0,0 +1,287 @@
|
||||
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
test/manual/test_torch_backend.py
Normal file
224
test/manual/test_torch_backend.py
Normal file
@@ -0,0 +1,224 @@
|
||||
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()
|
||||
Reference in New Issue
Block a user