File size: 3,733 Bytes
b88f761
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
"""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))