[Feature] add multi-rank support for Lora (#4492)

Co-authored-by: rudy152 <czh1137892874@gmail.com>
This commit is contained in:
chaobo jia
2025-03-28 09:38:44 -07:00
committed by GitHub
co-authored by rudy152
parent 6dea5c96bf
commit ef9a378a20
16 changed files with 292 additions and 97 deletions
+19 -33
View File
@@ -23,14 +23,10 @@ class BaseLayerWithLoRA(nn.Module):
def __init__(
self,
base_layer: nn.Module,
lora_rank: int,
scaling: float,
lora_backend: BaseLoRABackend,
):
super().__init__()
self.base_layer: nn.Module = base_layer
self.lora_rank: int = lora_rank
self.scaling: float = scaling
self.set_lora: bool = False
self.lora_backend: BaseLoRABackend = lora_backend
@@ -59,11 +55,9 @@ class VocabParallelEmbeddingWithLoRA(BaseLayerWithLoRA):
def __init__(
self,
base_layer: VocabParallelEmbedding,
lora_rank: int,
scaling: float,
lora_backend: BaseLoRABackend,
) -> None:
super().__init__(base_layer, lora_rank, scaling, lora_backend)
super().__init__(base_layer, lora_backend)
self.weight = base_layer.weight
@@ -71,11 +65,9 @@ class ColumnParallelLinearWithLoRA(BaseLayerWithLoRA):
def __init__(
self,
base_layer: ColumnParallelLinear,
lora_rank: int,
scaling: float,
lora_backend: BaseLoRABackend,
) -> None:
super().__init__(base_layer, lora_rank, scaling, lora_backend)
super().__init__(base_layer, lora_backend)
def set_lora_info(
self,
@@ -87,7 +79,7 @@ class ColumnParallelLinearWithLoRA(BaseLayerWithLoRA):
self.B_buffer = B_buffer
def apply_lora(self, base_output: torch.Tensor, x: torch.Tensor) -> torch.Tensor:
backend_kwargs = {"base_output": base_output, "scaling": self.scaling}
backend_kwargs = {"base_output": base_output}
lora_a_output = self.lora_backend.run_lora_a_sgemm(x, self.A_buffer)
lora_output = self.lora_backend.run_lora_b_sgemm(
lora_a_output,
@@ -96,8 +88,8 @@ class ColumnParallelLinearWithLoRA(BaseLayerWithLoRA):
)
return (
lora_output
if self.lora_backend.fuse_output_scaling_add
else base_output + lora_output * self.scaling
if self.lora_backend.fuse_output_add
else base_output + lora_output
)
def forward(self, input_: torch.Tensor):
@@ -132,11 +124,9 @@ class MergedColumnParallelLinearWithLoRA(ColumnParallelLinearWithLoRA):
def __init__(
self,
base_layer: MergedColumnParallelLinear,
lora_rank: int,
scaling: float,
lora_backend: BaseLoRABackend,
) -> None:
super().__init__(base_layer, lora_rank, scaling, lora_backend)
super().__init__(base_layer, lora_backend)
def set_lora_info(
self,
@@ -155,7 +145,7 @@ class MergedColumnParallelLinearWithLoRA(ColumnParallelLinearWithLoRA):
self.B_buffer_gate_up = (B_buffer[0], B_buffer[1])
def apply_lora(self, base_output: torch.Tensor, x: torch.Tensor) -> torch.Tensor:
backend_kwargs = {"base_output": base_output, "scaling": self.scaling}
backend_kwargs = {"base_output": base_output}
lora_output = self.lora_backend.run_gate_up_lora(
x,
@@ -165,8 +155,8 @@ class MergedColumnParallelLinearWithLoRA(ColumnParallelLinearWithLoRA):
)
return (
lora_output
if self.lora_backend.fuse_output_scaling_add
else base_output + lora_output * self.scaling
if self.lora_backend.fuse_output_add
else base_output + lora_output
)
def slice_lora_a_weights(self, A: torch.Tensor, tp_rank: int):
@@ -184,11 +174,9 @@ class QKVParallelLinearWithLoRA(ColumnParallelLinearWithLoRA):
def init__(
self,
base_layer: QKVParallelLinear,
lora_rank: int,
scaling: float,
lora_backend: BaseLoRABackend,
) -> None:
super().__init__(base_layer, lora_rank, scaling, lora_backend)
super().__init__(base_layer, lora_backend)
def set_lora_info(
self,
@@ -230,7 +218,7 @@ class QKVParallelLinearWithLoRA(ColumnParallelLinearWithLoRA):
)
def apply_lora(self, base_output: torch.Tensor, x: torch.Tensor) -> torch.Tensor:
backend_kwargs = {"base_output": base_output, "scaling": self.scaling}
backend_kwargs = {"base_output": base_output}
if self.lora_backend.fuse_stacked_lora_b:
backend_kwargs["output_offset"] = self.output_offset
backend_kwargs["max_qkv_out_dim"] = self.max_qkv_out_dim
@@ -243,8 +231,8 @@ class QKVParallelLinearWithLoRA(ColumnParallelLinearWithLoRA):
)
return (
lora_output
if self.lora_backend.fuse_output_scaling_add
else base_output + lora_output * self.scaling
if self.lora_backend.fuse_output_add
else base_output + lora_output
)
def slice_lora_a_weights(self, A: torch.Tensor, tp_rank: int):
@@ -273,11 +261,9 @@ class RowParallelLinearWithLoRA(BaseLayerWithLoRA):
def __init__(
self,
base_layer: RowParallelLinear,
lora_rank: int,
scaling: float,
lora_backend: BaseLoRABackend,
) -> None:
super().__init__(base_layer, lora_rank, scaling, lora_backend)
super().__init__(base_layer, lora_backend)
def set_lora_info(self, A_buffer: torch.Tensor, B_buffer: torch.Tensor):
self.set_lora = True
@@ -285,7 +271,7 @@ class RowParallelLinearWithLoRA(BaseLayerWithLoRA):
self.B_buffer = B_buffer
def apply_lora(self, base_output: torch.Tensor, x: torch.Tensor) -> torch.Tensor:
backend_kwargs = {"base_output": base_output, "scaling": self.scaling}
backend_kwargs = {"base_output": base_output}
lora_a_output = self.lora_backend.run_lora_a_sgemm(x, self.A_buffer)
lora_output = self.lora_backend.run_lora_b_sgemm(
lora_a_output,
@@ -294,8 +280,8 @@ class RowParallelLinearWithLoRA(BaseLayerWithLoRA):
)
return (
lora_output
if self.lora_backend.fuse_output_scaling_add
else base_output + lora_output * self.scaling
if self.lora_backend.fuse_output_add
else base_output + lora_output
)
def forward(self, input_: torch.Tensor):
@@ -344,7 +330,7 @@ class RowParallelLinearWithLoRA(BaseLayerWithLoRA):
def get_lora_layer(
layer: nn.Module, lora_rank: int, scaling: int, lora_backend: BaseLoRABackend
layer: nn.Module, lora_backend: BaseLoRABackend
) -> BaseLayerWithLoRA:
supported_layer_types = {
# the order matters
@@ -356,6 +342,6 @@ def get_lora_layer(
}
for src_layer_type, lora_layer_type in supported_layer_types.items():
if isinstance(layer, src_layer_type): # pylint: disable=unidiomatic-typecheck
ret = lora_layer_type(layer, lora_rank, scaling, lora_backend)
ret = lora_layer_type(layer, lora_backend)
return ret
raise Exception(f"No corresponding LoRA layer supported for {type(layer)}.")