[diffusion] lora: Fix LoRA weight merging for torch.nn.Linear layers from diffusers modules (#14150)
Co-authored-by: niehen6174 <niehen.6174@gmail.com>
This commit is contained in:
co-authored by
niehen6174
parent
5ddd2f6b6b
commit
990023e59b
@@ -351,6 +351,46 @@ class RowParallelLinearWithLoRA(BaseLayerWithLoRA):
|
|||||||
return B
|
return B
|
||||||
|
|
||||||
|
|
||||||
|
class LinearWithLoRA(BaseLayerWithLoRA):
|
||||||
|
"""
|
||||||
|
Wrapper for standard torch.nn.Linear to support LoRA.
|
||||||
|
Unlike custom LinearBase classes, nn.Linear.forward() returns a single tensor,
|
||||||
|
not a tuple of (output, bias).
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
base_layer: nn.Linear,
|
||||||
|
lora_rank: int | None = None,
|
||||||
|
lora_alpha: int | None = None,
|
||||||
|
) -> None:
|
||||||
|
super().__init__(base_layer, lora_rank, lora_alpha)
|
||||||
|
|
||||||
|
@torch.compile()
|
||||||
|
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
|
lora_A = self.lora_A
|
||||||
|
lora_B = self.lora_B
|
||||||
|
if isinstance(self.lora_B, DTensor):
|
||||||
|
lora_B = self.lora_B.to_local()
|
||||||
|
lora_A = self.lora_A.to_local()
|
||||||
|
|
||||||
|
if not self.merged and not self.disable_lora:
|
||||||
|
lora_A_sliced = self.slice_lora_a_weights(lora_A.to(x, non_blocking=True))
|
||||||
|
lora_B_sliced = self.slice_lora_b_weights(lora_B.to(x, non_blocking=True))
|
||||||
|
delta = x @ lora_A_sliced.T @ lora_B_sliced.T
|
||||||
|
if self.lora_alpha != self.lora_rank:
|
||||||
|
delta = delta * (
|
||||||
|
self.lora_alpha / self.lora_rank # type: ignore
|
||||||
|
) # type: ignore
|
||||||
|
# nn.Linear.forward() returns a single tensor, not a tuple
|
||||||
|
out = self.base_layer(x)
|
||||||
|
return out + delta
|
||||||
|
else:
|
||||||
|
# nn.Linear.forward() returns a single tensor
|
||||||
|
out = self.base_layer(x)
|
||||||
|
return out
|
||||||
|
|
||||||
|
|
||||||
def wrap_with_lora_layer(
|
def wrap_with_lora_layer(
|
||||||
layer: nn.Module,
|
layer: nn.Module,
|
||||||
lora_rank: int | None = None,
|
lora_rank: int | None = None,
|
||||||
@@ -359,7 +399,9 @@ def wrap_with_lora_layer(
|
|||||||
"""
|
"""
|
||||||
transform the given layer to its corresponding LoRA layer
|
transform the given layer to its corresponding LoRA layer
|
||||||
"""
|
"""
|
||||||
supported_layer_types: dict[type[LinearBase], type[BaseLayerWithLoRA]] = {
|
supported_layer_types: dict[
|
||||||
|
type[LinearBase] | type[nn.Linear], type[BaseLayerWithLoRA]
|
||||||
|
] = {
|
||||||
# the order matters
|
# the order matters
|
||||||
# VocabParallelEmbedding: VocabParallelEmbeddingWithLoRA,
|
# VocabParallelEmbedding: VocabParallelEmbeddingWithLoRA,
|
||||||
QKVParallelLinear: QKVParallelLinearWithLoRA,
|
QKVParallelLinear: QKVParallelLinearWithLoRA,
|
||||||
@@ -367,9 +409,10 @@ def wrap_with_lora_layer(
|
|||||||
ColumnParallelLinear: ColumnParallelLinearWithLoRA,
|
ColumnParallelLinear: ColumnParallelLinearWithLoRA,
|
||||||
RowParallelLinear: RowParallelLinearWithLoRA,
|
RowParallelLinear: RowParallelLinearWithLoRA,
|
||||||
ReplicatedLinear: BaseLayerWithLoRA,
|
ReplicatedLinear: BaseLayerWithLoRA,
|
||||||
|
nn.Linear: LinearWithLoRA,
|
||||||
}
|
}
|
||||||
for src_layer_type, lora_layer_type in supported_layer_types.items():
|
for src_layer_type, lora_layer_type in supported_layer_types.items():
|
||||||
if isinstance(layer, src_layer_type): # pylint: disable=unidiomatic-typecheck
|
if isinstance(layer, src_layer_type): # type: ignore[arg-type]
|
||||||
ret = lora_layer_type(
|
ret = lora_layer_type(
|
||||||
layer,
|
layer,
|
||||||
lora_rank=lora_rank,
|
lora_rank=lora_rank,
|
||||||
|
|||||||
Reference in New Issue
Block a user