Files
BobRandomNumber-ComfyUI-Kyu…/moshi_src/moshi/modules/lora.py
T
2025-07-08 22:18:02 -04:00

123 lines
4.4 KiB
Python

import torch
import torch.nn as nn
def replace_all_linear_with_lora(module, rank: int, scaling: float, device=None, dtype=None):
""" Recursively replace all Linear layers with LoRALinear layers."""
for name, child in module.named_children():
if isinstance(child, nn.Linear):
if device is None:
this_device = child.weight.device
else:
this_device = device
if dtype is None:
this_dtype = child.weight.dtype
else:
this_dtype = dtype
lora = LoRALinear(child.in_features, child.out_features,
rank, scaling, device=this_device, dtype=this_dtype)
lora.frozen_W = child
setattr(module, name, lora)
else:
replace_all_linear_with_lora(child, rank, scaling, device=device, dtype=dtype)
def replace_lora_with_linear(module):
"""Recursively replace all LoRALinear layers with Linear layers."""
for name, child in module.named_children():
if isinstance(child, LoRALinear):
# Compute merged weights: W' = W + scaling * B @ A
merged_weight = child.frozen_W.weight.data + \
child.scaling * (child.lora_B.weight @ child.lora_A.weight)
# Create a standard Linear layer with the same in/out features
new_linear = nn.Linear(child.frozen_W.in_features,
child.frozen_W.out_features, bias=False,
device=torch.device('meta'),
dtype=merged_weight.dtype)
new_linear.weight = nn.Parameter(
merged_weight, requires_grad=merged_weight.requires_grad) # Transfer merged weights
setattr(module, name, new_linear) # Replace the module
else:
replace_lora_with_linear(child) # Recursively process submodules
class LoRALinear(nn.Module):
"""
Implementation of:
- LoRA: https://arxiv.org/abs/2106.09685
Notes:
- Freezing is handled at the network level, not the layer level.
- Scaling factor controls relative importance of LoRA skip
connection versus original frozen weight. General guidance is
to keep it to 2.0 and sweep over learning rate when changing
the rank.
"""
def __init__(
self,
in_features: int,
out_features: int,
rank: int,
scaling: float,
bias: bool = False,
device: torch.device | None = None,
dtype: torch.dtype = torch.bfloat16,
):
super().__init__()
self.in_features = in_features
self.out_features = out_features
assert not bias
self.bias = bias
self.rank = rank
self.scaling = scaling
self.lora_A = nn.Linear(
self.in_features,
self.rank,
bias=self.bias,
device=device,
dtype=dtype,
)
self.lora_B = nn.Linear(
self.rank,
self.out_features,
bias=self.bias,
device=device,
dtype=dtype,
)
self.frozen_W = nn.Linear(self.in_features,
self.out_features,
bias=self.bias,
device=device,
dtype=dtype)
self._register_load_state_dict_pre_hook(LoRALinear._load_hook, with_module=True)
def merge_weight(self):
with torch.no_grad():
down_weight = self.lora_A.weight
up_weight = self.lora_B.weight
weight = up_weight.mm(down_weight) * self.scaling
weight += self.frozen_W.weight
return weight
@staticmethod
def _load_hook(module, state_dict, prefix, *_):
key_name = prefix + "weight"
if key_name in state_dict:
w_ref = state_dict.pop(key_name)
state_dict[prefix + 'frozen_W.weight'] = w_ref
def forward(self, x: torch.Tensor):
lora = self.lora_B(self.lora_A(x))
return self.frozen_W(x) + lora * self.scaling
def __repr__(self) -> str:
return "{}Linear(in_features={}, out_features={}, r={})".format(
"LoRA", self.in_features, self.out_features, self.rank)