[1/2] Refactor LoRA to support backend-specific batch preprocessing. (#10251)
This commit is contained in:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user