TrinityX / models /mocae_layer.py
Gautam Kashyap
Upload folder using huggingface_hub
e38f140 verified
Raw
History Blame Contribute Delete
2.93 kB
import torch
import torch.nn as nn
import torch.nn.functional as F
from typing import Optional
from .mocae_router import MoCaERouter
class LoRAExpertFFN(nn.Module):
def __init__(self, base_ffn: nn.Module, lora_weights: dict, lora_scale: float = 2.0):
super().__init__()
self.base_ffn = base_ffn
self.lora_scale = lora_scale
for key, tensor in lora_weights.items():
self.register_buffer(key, tensor.contiguous())
def _lora(self, x: torch.Tensor, A_name: str, B_name: str) -> torch.Tensor:
A = getattr(self, A_name)
B = getattr(self, B_name)
return (x @ A.T.to(x.dtype)) @ B.T.to(x.dtype) * self.lora_scale
def forward(self, h: torch.Tensor) -> torch.Tensor:
gate = self.base_ffn.gate_proj(h) + self._lora(h, "gate_A", "gate_B")
up = self.base_ffn.up_proj(h) + self._lora(h, "up_A", "up_B")
intermediate = F.silu(gate) * up
out = self.base_ffn.down_proj(intermediate) + self._lora(intermediate, "down_A", "down_B")
return out
class MoCaELayer(nn.Module):
def __init__(
self,
original_ffn: nn.Module,
expert_ffns: list,
router: MoCaERouter,
gamma_tilde: list,
hidden_size: int,
dropout: float = 0.1,
):
super().__init__()
self.expert_ffns = nn.ModuleList(expert_ffns)
self.router = router
self.gamma_tilde = gamma_tilde
self.layer_norm = nn.LayerNorm(hidden_size)
self.dropout = nn.Dropout(p=dropout)
self.prev_probs: Optional[torch.Tensor] = None
def forward(self, h: torch.Tensor) -> torch.Tensor:
orig_dtype = h.dtype
router_device = next(self.router.parameters()).device
h = h.to(device=router_device, dtype=torch.bfloat16)
probs, _ = self.router(h)
expert_outputs = torch.stack(
[ffn(h) for ffn in self.expert_ffns], dim=-1
)
gamma = torch.tensor(self.gamma_tilde, device=h.device, dtype=h.dtype)
alpha = probs * gamma.unsqueeze(0).unsqueeze(0)
alpha = alpha / (alpha.sum(dim=-1, keepdim=True) + 1e-9)
y = (expert_outputs * alpha.unsqueeze(-2)).sum(dim=-1)
y_cal = self.dropout(y)
self.prev_probs = probs.detach()
return y_cal.to(orig_dtype)
def get_prev_probs(self):
return self.prev_probs
def update_gamma(self, new_gamma: list):
self.gamma_tilde = new_gamma
def get_expert_activations(self, h: torch.Tensor) -> dict:
probs, _ = self.router(h)
gamma = torch.tensor(self.gamma_tilde, device=h.device, dtype=h.dtype)
alpha = probs * gamma.unsqueeze(0).unsqueeze(0)
alpha = alpha / (alpha.sum(dim=-1, keepdim=True) + 1e-9)
return {
"routing_probs": probs.mean(dim=(0, 1)).cpu().tolist(),
"combined_weights": alpha.mean(dim=(0, 1)).cpu().tolist(),
}