[1/2] Refactor LoRA to support backend-specific batch preprocessing. (#10251)

This commit is contained in:
Lifu Huang
2025-09-10 09:58:37 -07:00
committed by GitHub
parent cda7e47ce7
commit 941002945b
6 changed files with 227 additions and 130 deletions
+32
View File
@@ -66,6 +66,15 @@ class ColumnParallelLinearWithLoRA(BaseLayerWithLoRA):
lora_backend: BaseLoRABackend,
) -> None:
super().__init__(base_layer, lora_backend)
shard_size = self.base_layer.output_partition_sizes[0]
self.output_offset = torch.tensor(
[
0,
shard_size,
],
dtype=torch.int32,
device=next(self.base_layer.parameters()).device,
)
def set_lora_info(
self,
@@ -81,6 +90,7 @@ class ColumnParallelLinearWithLoRA(BaseLayerWithLoRA):
lora_output = self.lora_backend.run_lora_b_sgemm(
x=lora_a_output,
weights=self.B_buffer,
output_offset=self.output_offset,
base_output=base_output,
)
return lora_output
@@ -130,11 +140,23 @@ class MergedColumnParallelLinearWithLoRA(ColumnParallelLinearWithLoRA):
self.A_buffer_gate_up = A_buffer
self.B_buffer_gate_up = B_buffer
shard_size = self.base_layer.output_partition_sizes[0]
self.output_offset = torch.tensor(
[
0,
shard_size,
2 * shard_size,
],
dtype=torch.int32,
device=next(self.base_layer.parameters()).device,
)
def apply_lora(self, base_output: torch.Tensor, x: torch.Tensor) -> torch.Tensor:
lora_output = self.lora_backend.run_gate_up_lora(
x=x,
gate_up_lora_a=self.A_buffer_gate_up,
gate_up_lora_b=self.B_buffer_gate_up,
output_offset=self.output_offset,
base_output=base_output,
)
return lora_output
@@ -243,12 +265,22 @@ class RowParallelLinearWithLoRA(BaseLayerWithLoRA):
self.set_lora = True
self.A_buffer = A_buffer
self.B_buffer = B_buffer
output_size = self.base_layer.output_size
self.output_offset = torch.tensor(
[
0,
output_size,
],
dtype=torch.int32,
device=next(self.base_layer.parameters()).device,
)
def apply_lora(self, base_output: torch.Tensor, x: torch.Tensor) -> torch.Tensor:
lora_a_output = self.lora_backend.run_lora_a_sgemm(x, self.A_buffer)
lora_output = self.lora_backend.run_lora_b_sgemm(
x=lora_a_output,
weights=self.B_buffer,
output_offset=self.output_offset,
base_output=base_output,
)
return lora_output