TrinityX / models /mocae_model.py
Gautam Kashyap
Upload folder using huggingface_hub
e38f140 verified
Raw
History Blame Contribute Delete
5.93 kB
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