File size: 2,933 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
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(),
        }