LoRA: Ascend: Update ascend LoRA backend to support new kernels (#15912)
This commit is contained in:
@@ -29,22 +29,17 @@ class AscendLoRABackend(BaseLoRABackend):
|
||||
_, weight_out_dim, _ = weights.shape
|
||||
|
||||
output_tensor = torch.zeros(
|
||||
(total_seq_len, weight_out_dim), dtype=x.dtype, device=x.device
|
||||
(total_seq_len, weight_out_dim), dtype=torch.float, device=x.device
|
||||
)
|
||||
torch.ops.npu.sgmv_shrink(
|
||||
torch.ops.npu.sgemmv_shrink(
|
||||
x,
|
||||
weights,
|
||||
self.batch_info.weight_indices,
|
||||
self.batch_info.seg_lens,
|
||||
self.batch_info.lora_ranks,
|
||||
self.batch_info.scalings,
|
||||
output_tensor,
|
||||
1.0,
|
||||
)
|
||||
scaling = (
|
||||
self.batch_info.scalings.gather(0, self.batch_info.weight_indices)
|
||||
.repeat_interleave(self.batch_info.seg_lens, output_size=total_seq_len)
|
||||
.unsqueeze(-1)
|
||||
)
|
||||
output_tensor *= scaling
|
||||
|
||||
return output_tensor
|
||||
|
||||
@@ -52,6 +47,7 @@ class AscendLoRABackend(BaseLoRABackend):
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
weights: torch.Tensor,
|
||||
output_offset: torch.Tensor,
|
||||
base_output: torch.Tensor = None,
|
||||
*args,
|
||||
**kwargs,
|
||||
@@ -66,14 +62,14 @@ class AscendLoRABackend(BaseLoRABackend):
|
||||
else:
|
||||
output_tensor = base_output
|
||||
|
||||
torch.ops.npu.sgmv_expand(
|
||||
torch.ops.npu.sgemmv_expand(
|
||||
x,
|
||||
weights,
|
||||
self.batch_info.weight_indices,
|
||||
self.batch_info.seg_lens,
|
||||
self.batch_info.lora_ranks,
|
||||
output_offset,
|
||||
output_tensor,
|
||||
0,
|
||||
weight_out_dim,
|
||||
)
|
||||
|
||||
return output_tensor
|
||||
@@ -84,8 +80,6 @@ class AscendLoRABackend(BaseLoRABackend):
|
||||
qkv_lora_a: torch.Tensor,
|
||||
qkv_lora_b: torch.Tensor,
|
||||
output_offset: torch.Tensor,
|
||||
output_offset_cpu: torch.Tensor,
|
||||
max_qkv_out_dim: int,
|
||||
base_output: torch.Tensor = None,
|
||||
n_slices: int = 3,
|
||||
*args,
|
||||
@@ -106,37 +100,28 @@ class AscendLoRABackend(BaseLoRABackend):
|
||||
output_tensor = base_output
|
||||
|
||||
lora_a_output = torch.zeros(
|
||||
total_seq_len, weight_intermediate_dim, dtype=x.dtype, device=x.device
|
||||
total_seq_len, weight_intermediate_dim, dtype=torch.float, device=x.device
|
||||
)
|
||||
torch.ops.npu.sgmv_shrink(
|
||||
|
||||
torch.ops.npu.sgemmv_shrink(
|
||||
x,
|
||||
qkv_lora_a,
|
||||
self.batch_info.weight_indices,
|
||||
self.batch_info.seg_lens,
|
||||
self.batch_info.lora_ranks,
|
||||
self.batch_info.scalings,
|
||||
lora_a_output,
|
||||
1.0,
|
||||
)
|
||||
|
||||
scaling = (
|
||||
self.batch_info.scalings.gather(0, self.batch_info.weight_indices)
|
||||
.repeat_interleave(self.batch_info.seg_lens, output_size=total_seq_len)
|
||||
.unsqueeze(-1)
|
||||
torch.ops.npu.sgemmv_expand(
|
||||
lora_a_output,
|
||||
qkv_lora_b,
|
||||
self.batch_info.weight_indices,
|
||||
self.batch_info.seg_lens,
|
||||
self.batch_info.lora_ranks,
|
||||
output_offset,
|
||||
output_tensor,
|
||||
)
|
||||
lora_a_output *= scaling
|
||||
|
||||
for slice_id in range(n_slices):
|
||||
slice_offset = output_offset_cpu[slice_id]
|
||||
slice_offset_next = output_offset_cpu[slice_id + 1]
|
||||
slice_size = slice_offset_next - slice_offset
|
||||
torch.ops.npu.sgmv_expand(
|
||||
lora_a_output[:, (max_rank * slice_id) : (max_rank * (slice_id + 1))],
|
||||
qkv_lora_b[:, slice_offset:slice_offset_next],
|
||||
self.batch_info.weight_indices,
|
||||
self.batch_info.seg_lens,
|
||||
output_tensor,
|
||||
slice_offset,
|
||||
slice_size,
|
||||
)
|
||||
|
||||
return output_tensor
|
||||
|
||||
@@ -145,19 +130,17 @@ class AscendLoRABackend(BaseLoRABackend):
|
||||
x: torch.Tensor,
|
||||
gate_up_lora_a: torch.Tensor,
|
||||
gate_up_lora_b: torch.Tensor,
|
||||
output_offset: torch.Tensor,
|
||||
base_output: torch.Tensor = None,
|
||||
*args,
|
||||
**kwargs,
|
||||
) -> torch.Tensor:
|
||||
|
||||
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
|
||||
slice_size = weight_out_dim // num_slices
|
||||
max_rank = weight_intermediate_dim // num_slices
|
||||
|
||||
if base_output is None:
|
||||
output_tensor = torch.zeros(
|
||||
@@ -167,37 +150,28 @@ class AscendLoRABackend(BaseLoRABackend):
|
||||
output_tensor = base_output
|
||||
|
||||
lora_a_output = torch.zeros(
|
||||
total_seq_len, weight_intermediate_dim, dtype=x.dtype, device=x.device
|
||||
total_seq_len, weight_intermediate_dim, dtype=torch.float, device=x.device
|
||||
)
|
||||
|
||||
torch.ops.npu.sgmv_shrink(
|
||||
torch.ops.npu.sgemmv_shrink(
|
||||
x,
|
||||
gate_up_lora_a,
|
||||
self.batch_info.weight_indices,
|
||||
self.batch_info.seg_lens,
|
||||
self.batch_info.lora_ranks,
|
||||
self.batch_info.scalings,
|
||||
lora_a_output,
|
||||
1.0,
|
||||
)
|
||||
|
||||
scaling = (
|
||||
self.batch_info.scalings.gather(0, self.batch_info.weight_indices)
|
||||
.repeat_interleave(self.batch_info.seg_lens, output_size=total_seq_len)
|
||||
.unsqueeze(-1)
|
||||
torch.ops.npu.sgemmv_expand(
|
||||
lora_a_output,
|
||||
gate_up_lora_b,
|
||||
self.batch_info.weight_indices,
|
||||
self.batch_info.seg_lens,
|
||||
self.batch_info.lora_ranks,
|
||||
output_offset,
|
||||
output_tensor,
|
||||
)
|
||||
lora_a_output *= scaling
|
||||
|
||||
slice_offset = 0
|
||||
for slice_id in range(num_slices):
|
||||
torch.ops.npu.sgmv_expand(
|
||||
lora_a_output[:, (max_rank * slice_id) : (max_rank * (slice_id + 1))],
|
||||
gate_up_lora_b[:, slice_offset : slice_offset + slice_size],
|
||||
self.batch_info.weight_indices,
|
||||
self.batch_info.seg_lens,
|
||||
output_tensor,
|
||||
slice_offset,
|
||||
slice_size,
|
||||
)
|
||||
slice_offset += slice_size
|
||||
|
||||
return output_tensor
|
||||
|
||||
@@ -218,7 +192,7 @@ class AscendLoRABackend(BaseLoRABackend):
|
||||
max_len=num_tokens_per_req,
|
||||
weight_indices=torch.zeros(max_bs_in_cuda_graph, 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.float16),
|
||||
permutation=None,
|
||||
)
|
||||
|
||||
@@ -246,7 +220,7 @@ class AscendLoRABackend(BaseLoRABackend):
|
||||
lora_ranks, dtype=torch.int32, pin_memory=True, device="cpu"
|
||||
)
|
||||
scalings_tensor = torch.tensor(
|
||||
scalings, dtype=torch.float, pin_memory=True, device="cpu"
|
||||
scalings, dtype=torch.float16, pin_memory=True, device="cpu"
|
||||
)
|
||||
|
||||
bs = forward_batch.batch_size
|
||||
@@ -287,7 +261,7 @@ class AscendLoRABackend(BaseLoRABackend):
|
||||
(self.max_loras_per_batch,), dtype=torch.int32, device=self.device
|
||||
),
|
||||
scalings=torch.empty(
|
||||
(self.max_loras_per_batch,), dtype=torch.float, device=self.device
|
||||
(self.max_loras_per_batch,), dtype=torch.float16, device=self.device
|
||||
),
|
||||
permutation=None,
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user