[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:
@@ -351,6 +351,46 @@ class RowParallelLinearWithLoRA(BaseLayerWithLoRA):
|
||||
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(
|
||||
layer: nn.Module,
|
||||
lora_rank: int | None = None,
|
||||
@@ -359,7 +399,9 @@ def wrap_with_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
|
||||
# VocabParallelEmbedding: VocabParallelEmbeddingWithLoRA,
|
||||
QKVParallelLinear: QKVParallelLinearWithLoRA,
|
||||
@@ -367,9 +409,10 @@ def wrap_with_lora_layer(
|
||||
ColumnParallelLinear: ColumnParallelLinearWithLoRA,
|
||||
RowParallelLinear: RowParallelLinearWithLoRA,
|
||||
ReplicatedLinear: BaseLayerWithLoRA,
|
||||
nn.Linear: LinearWithLoRA,
|
||||
}
|
||||
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(
|
||||
layer,
|
||||
lora_rank=lora_rank,
|
||||
|
||||
Reference in New Issue
Block a user