File size: 5,929 Bytes
e38f140
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
import os
import json
import torch
import torch.nn as nn
from typing import Optional
from safetensors.torch import load_file

from .mocae_router import MoCaERouter
from .mocae_layer import MoCaELayer, LoRAExpertFFN


def load_lora_weights(adapter_path: str) -> dict:
    weights = load_file(adapter_path)
    return {k: v.cpu() for k, v in weights.items()}


def get_ffn_lora_for_layer(lora_weights: dict, layer_idx: int, lora_alpha: int = 32, lora_rank: int = 16) -> dict:
    scale = lora_alpha / lora_rank
    prefix = f"base_model.model.model.layers.{layer_idx}.mlp"

    def get(module, ab):
        key = f"{prefix}.{module}.lora_{ab}.weight"
        return lora_weights.get(key)

    result = {}
    for name, module in [("gate", "gate_proj"), ("up", "up_proj"), ("down", "down_proj")]:
        A = get(module, "A")
        B = get(module, "B")
        if A is not None and B is not None:
            result[f"{name}_A"] = A.to(torch.bfloat16)
            result[f"{name}_B"] = B.to(torch.bfloat16)

    return result


def compute_gamma_from_lora(adapter_paths: list) -> list:
    all_weights = [load_lora_weights(os.path.join(p, "adapter_model.safetensors")) for p in adapter_paths]

    ref = {}
    n = len(all_weights)
    for w_dict in all_weights:
        for k, v in w_dict.items():
            if k not in ref:
                ref[k] = v.float() / n
            else:
                ref[k] = ref[k] + v.float() / n

    raw = []
    for w_dict in all_weights:
        ip = sum(
            (w_dict[k].float() * ref[k]).sum().item()
            for k in w_dict if k in ref
        )
        raw.append(abs(ip) + 1e-9)

    total = sum(raw)
    gamma = [r / total for r in raw]
    return gamma


class TrinityXModel(nn.Module):

    def __init__(
        self,
        base_model,
        adapter_paths: list,
        gamma_tilde: list,
        mocae_config: dict,
        lora_rank: int = 16,
        lora_alpha: int = 32,
    ):
        super().__init__()
        self.base_model = base_model
        self.gamma_tilde = gamma_tilde
        self.mocae_config = mocae_config
        self.mocae_layers: list = []

        hidden_size = base_model.config.hidden_size
        self._inject_mocae_layers(adapter_paths, hidden_size, lora_rank, lora_alpha)

    def _inject_mocae_layers(self, adapter_paths, hidden_size, lora_rank, lora_alpha):
        print(f"[TrinityX] Loading {len(adapter_paths)} LoRA adapters...")
        all_lora = [load_lora_weights(os.path.join(p, "adapter_model.safetensors"))
                    for p in adapter_paths]

        for layer_idx, layer in enumerate(self.base_model.model.layers):
            original_ffn = layer.mlp

            expert_ffns = []
            for lora_weights in all_lora:
                ffn_lora = get_ffn_lora_for_layer(lora_weights, layer_idx, lora_alpha, lora_rank)
                if ffn_lora:
                    expert_ffns.append(LoRAExpertFFN(
                        base_ffn=original_ffn,
                        lora_weights=ffn_lora,
                        lora_scale=lora_alpha / lora_rank,
                    ))
                else:
                    expert_ffns.append(original_ffn)

            router = MoCaERouter(
                hidden_size=hidden_size,
                num_experts=self.mocae_config.get("num_experts", 3),
                temperature=self.mocae_config.get("temperature", 0.7),
                router_hidden_dim=self.mocae_config.get("router_hidden_dim", 128),
                router_output_dim=self.mocae_config.get("router_output_dim", 64),
            )

            mocae = MoCaELayer(
                original_ffn=original_ffn,
                expert_ffns=expert_ffns,
                router=router,
                gamma_tilde=list(self.gamma_tilde),
                hidden_size=hidden_size,
                dropout=self.mocae_config.get("dropout_rate", 0.1),
            )

            try:
                layer_dev = next(p.device for name, p in layer.named_parameters()
                                 if 'self_attn' in name or 'input_layernorm' in name)
            except StopIteration:
                layer_dev = torch.device("cuda:0")

            mocae.router.to(device=layer_dev, dtype=torch.bfloat16)
            mocae.layer_norm.to(device=layer_dev, dtype=torch.bfloat16)

            for exp_ffn in mocae.expert_ffns:
                for buf_name in list(exp_ffn._buffers.keys()):
                    exp_ffn._buffers[buf_name] = exp_ffn._buffers[buf_name].to(
                        device=layer_dev, dtype=torch.bfloat16)

            layer.mlp = mocae
            self.mocae_layers.append(mocae)

        print(f"[TrinityX] MoCaE injected at {len(self.mocae_layers)} layers")

    def forward(self, input_ids, attention_mask=None, labels=None, **kwargs):
        return self.base_model(
            input_ids=input_ids,
            attention_mask=attention_mask,
            labels=labels,
            **kwargs,
        )

    def update_all_gamma(self, new_gamma: list):
        self.gamma_tilde = new_gamma
        for layer in self.mocae_layers:
            layer.update_gamma(new_gamma)

    @classmethod
    def load_pretrained(cls, save_dir, base_model, adapter_paths, mocae_config,
                        lora_rank: int = 16, lora_alpha: int = 32):
        with open(os.path.join(save_dir, "gamma_tilde.json")) as f:
            gamma_tilde = json.load(f)["gamma_tilde"]
        model = cls(base_model, adapter_paths, gamma_tilde, mocae_config,
                    lora_rank=lora_rank, lora_alpha=lora_alpha)
        router_states = torch.load(os.path.join(save_dir, "mocae_routers.pt"), map_location="cpu")
        for i, layer in enumerate(model.mocae_layers):
            state = {k.replace(f"layer_{i}.router.", ""): v
                     for k, v in router_states.items() if k.startswith(f"layer_{i}.router.")}
            layer.router.load_state_dict(state)
        return model