"""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))