diff --git a/python/sglang/srt/lora/backend/ascend_backend.py b/python/sglang/srt/lora/backend/ascend_backend.py index 77924752f..8141f1354 100644 --- a/python/sglang/srt/lora/backend/ascend_backend.py +++ b/python/sglang/srt/lora/backend/ascend_backend.py @@ -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, )