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