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