ydy9038074's picture
Publish Modilify Mk2 Preview step 1250 schema25
b88f761 verified
Raw History Blame Contribute Delete
3.73 kB
"""Unfused dense and expert inference adapters."""
from __future__ import annotations
from dataclasses import dataclass
import math
import torch
from torch import nn
from transformers.models.diffusion_gemma.modeling_diffusion_gemma import DiffusionGemmaTextExperts
from .mps_ops import attach_expert_lora
@dataclass
class LoRAConfig:
r: int
alpha: int
target_modules: tuple[str, ...]
expert_r: int = 8
expert_alpha: int = 8
class LoRALinear(nn.Module):
def __init__(self, base: nn.Linear, config: LoRAConfig):
super().__init__()
if config.r <= 0:
raise ValueError("LoRA rank must be positive.")
self.base = base
self.scaling = config.alpha / config.r
self.lora_a = nn.Parameter(
torch.empty(config.r, base.in_features, device=base.weight.device, dtype=base.weight.dtype)
)
self.lora_b = nn.Parameter(
torch.zeros(base.out_features, config.r, device=base.weight.device, dtype=base.weight.dtype)
)
nn.init.kaiming_uniform_(self.lora_a, a=math.sqrt(5))
self.base.requires_grad_(False)
def forward(self, inputs: torch.Tensor) -> torch.Tensor:
update = torch.nn.functional.linear(
torch.nn.functional.linear(inputs, self.lora_a), self.lora_b,
)
return self.base(inputs) + self.scaling * update
def get_parent_module(root: nn.Module, module_name: str) -> tuple[nn.Module, str]:
parent = root
parts = module_name.split(".")
for part in parts[:-1]:
parent = getattr(parent, part)
return parent, parts[-1]
def inject_lora(model: nn.Module, config: LoRAConfig) -> int:
targets = set(config.target_modules)
replacements = []
for name, module in model.named_modules():
if not isinstance(module, nn.Linear):
continue
if ".router." in name:
continue
in_trunk = ".layers." in name
if name.rsplit(".", 1)[-1] not in targets or not in_trunk:
continue
replacements.append((name, module))
for name, module in replacements:
parent, child = get_parent_module(model, name)
setattr(parent, child, LoRALinear(module, config))
expert_modules = [
module for module in model.modules()
if isinstance(module, DiffusionGemmaTextExperts)
]
for module in expert_modules:
attach_expert_lora(
module,
rank=config.expert_r,
alpha=config.expert_alpha,
)
share_tied_lora_parameters(model)
if not replacements and not expert_modules:
raise RuntimeError("No LoRA target modules were found.")
return len(replacements) + len(expert_modules)
def share_tied_lora_parameters(model: nn.Module) -> None:
"""Use one adapter for the encoder/decoder pair whose base weights are tied."""
trunk = getattr(model, "model", model)
encoder = trunk.encoder.language_model
decoder = trunk.decoder
expert_names = ("lora_gate_up_a", "lora_gate_up_b", "lora_down_a", "lora_down_b")
for name, decoder_module in decoder.named_modules():
if not name:
continue
try:
encoder_module = encoder.get_submodule(name)
except AttributeError:
continue
if isinstance(decoder_module, LoRALinear) and isinstance(encoder_module, LoRALinear):
encoder_module.lora_a = decoder_module.lora_a
encoder_module.lora_b = decoder_module.lora_b
for parameter_name in expert_names:
if hasattr(decoder_module, parameter_name) and hasattr(encoder_module, parameter_name):
setattr(encoder_module, parameter_name, getattr(decoder_module, parameter_name))