File size: 7,560 Bytes
fec7158
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
"""

Unified Sparse-AST / BWM Architecture & Universal Subfolder Loader

Domain: Blender 3D Mathematics & Blender Python Libraries (bpy, mathutils, bmesh, numpy, gpu)

"""

import os
import json
import torch
import torch.nn as nn
import torch.nn.functional as F
import safetensors.torch as st

class RMSNorm(nn.Module):
    def __init__(self, d):
        super().__init__()
        self.w = nn.Parameter(torch.ones(d))
    def forward(self, x):
        return x * torch.rsqrt(x.float().square().mean(-1, keepdim=True) + 1e-6).to(x.dtype) * self.w

class SparseASTBlock(nn.Module):
    def __init__(self, d=800, h=1600, al_hidden=24, n_heads=8):
        super().__init__()
        self.n1 = RMSNorm(d)
        self.n2 = RMSNorm(d)
        self.a = nn.MultiheadAttention(d, n_heads, batch_first=True)
        self.up = nn.Linear(d, h)
        self.r = nn.Linear(h, 3)
        self.down = nn.Linear(2*h, d)
        self.al = nn.Sequential(nn.Linear(3*h + 3, al_hidden), nn.SiLU(), nn.Linear(al_hidden, 1))
        self.m = nn.Linear(d, d, bias=False)
        
    def forward(self, x):
        t = x.shape[1]
        q = self.n1(x)
        mask = torch.ones(t, t, device=x.device, dtype=torch.bool).triu(1)
        x = torch.clamp(x + self.a(q, q, q, attn_mask=mask, need_weights=False)[0], -50, 50)
        z = torch.clamp(self.up(self.n2(x)), -12, 12)
        g = F.gumbel_softmax(self.r(z), tau=2.0, hard=True, dim=-1)
        e = z.sign() * torch.expm1(z.abs().clamp(max=2.0))
        u = (torch.stack((z, z.sign() * torch.log1p(z.abs()), e), -1) * g.unsqueeze(-2)).sum(-1).clamp(-8, 8)
        p, n = F.relu(u), F.relu(-u)
        d = torch.exp(-F.softplus(self.al(torch.cat((p, n, u, g), -1))).squeeze(-1)).clamp(1e-4, 0.9999)
        y = torch.clamp(x + self.down(torch.cat((p, n), -1)), -50, 50)
        state = torch.zeros_like(y[:, 0])
        o = []
        for j in range(t):
            state = torch.clamp(state * d[:, j:j+1] + y[:, j] * (1 - d[:, j:j+1]), -50, 50)
            o.append(torch.clamp(y[:, j] + self.m(state), -50, 50))
        return torch.stack(o, 1)

class SparseAST(nn.Module):
    def __init__(self, d=800, h=1600, layers=28, seq_len=32, vocab=512, al_hidden=24, n_heads=8):
        super().__init__()
        self.d = d
        self.h = h
        self.layers = layers
        self.seq_len = seq_len
        self.vocab = vocab
        
        self.e = nn.Embedding(vocab, d)
        self.p = nn.Embedding(seq_len, d)
        self.b = nn.ModuleList([SparseASTBlock(d=d, h=h, al_hidden=al_hidden, n_heads=n_heads) for _ in range(layers)])
        self.n = RMSNorm(d)
        self.h_out = nn.Linear(d, vocab, bias=False)
        self.h_out.weight = self.e.weight  # Tied word embeddings

    def forward(self, i):
        if i.shape[1] > self.seq_len:
            i = i[:, -self.seq_len:]
        seq = i.shape[1]
        pos = torch.arange(seq, device=i.device)[None]
        x = self.e(i) + self.p(pos)
        for b in self.b:
            x = b(x)
        return self.h_out(self.n(x))

    @classmethod
    def from_pretrained(cls, repo_or_dir, subfolder="100M-32", device="cpu", target_context=None):
        target_dir = os.path.join(repo_or_dir, subfolder) if subfolder else repo_or_dir
        config_path = os.path.join(target_dir, "config.json")
        with open(config_path, "r", encoding="utf-8") as f:
            cfg = json.load(f)
            
        native_seq = cfg.get("native_seq_len", 32)
        effective_seq = max(native_seq, target_context) if target_context else native_seq
        
        model = cls(
            d=cfg["d_model"],
            h=cfg["d_hidden"],
            layers=cfg["num_layers"],
            seq_len=effective_seq,
            vocab=cfg["vocab_size"],
            al_hidden=cfg.get("al_hidden", 24),
            n_heads=cfg.get("num_attention_heads", 8)
        )
        
        weights_path = os.path.join(target_dir, "model.safetensors")
        state_dict = st.load_file(weights_path, device="cpu")
        
        if "h.weight" in state_dict and "h_out.weight" not in state_dict:
            state_dict["h_out.weight"] = state_dict["h.weight"]
        elif "h_out.weight" not in state_dict and "e.weight" in state_dict:
            state_dict["h_out.weight"] = state_dict["e.weight"]
            
        if effective_seq > native_seq and "p.weight" in state_dict:
            old_p = state_dict["p.weight"]
            new_p = F.interpolate(
                old_p.T.unsqueeze(0),
                size=effective_seq,
                mode="linear",
                align_corners=True
            ).squeeze(0).T
            state_dict["p.weight"] = new_p
            
        model.load_state_dict(state_dict)
        model.to(device)
        model.eval()
        return model

class TopKRouter(nn.Module):
    def __init__(self, vocab_size=512, hidden_dim=64, num_experts=4):
        super().__init__()
        self.embed = nn.Embedding(vocab_size, hidden_dim)
        self.norm = nn.LayerNorm(hidden_dim)
        self.mlp = nn.Sequential(
            nn.Linear(hidden_dim, hidden_dim),
            nn.SiLU(),
            nn.Linear(hidden_dim, num_experts)
        )
    def forward(self, x):
        emb = self.norm(self.embed(x))
        return self.mlp(emb)

class TopKSparseASTEnsemble(nn.Module):
    def __init__(self, router, experts=None, expert_names=None, k=2, tau=1.0, device="cpu"):
        super().__init__()
        self.router = router
        self.experts = nn.ModuleList(experts if experts else [])
        self.expert_names = expert_names or []
        self.k = k
        self.tau = tau
        self.device = device
        
    @classmethod
    def from_pretrained(cls, repo_or_dir, subfolder="TopK-MoE", device="cpu", load_experts=True):
        target_dir = os.path.join(repo_or_dir, subfolder) if subfolder else repo_or_dir
        cfg_path = os.path.join(target_dir, "config.json")
        with open(cfg_path, "r", encoding="utf-8") as f:
            cfg = json.load(f)
            
        router = TopKRouter(
            vocab_size=cfg["vocab_size"],
            hidden_dim=cfg["hidden_dim"],
            num_experts=cfg["num_experts"]
        )
        st_path = os.path.join(target_dir, "model.safetensors")
        router.load_state_dict(st.load_file(st_path, device=device))
        router.to(device)
        router.eval()
        
        experts = []
        names = cfg.get("expert_names", ["3M-32", "10M-32", "100M-32", "200M-32"])
        if load_experts:
            for exp_name in names:
                exp_sub = exp_name if exp_name.endswith("-32") else f"{exp_name}-32"
                exp_path = os.path.join(repo_or_dir, exp_sub)
                if os.path.exists(exp_path):
                    print(f"Loading MoE expert [{exp_name}] from {exp_sub}...")
                    m = SparseAST.from_pretrained(repo_or_dir, subfolder=exp_sub, device=device)
                    experts.append(m)
                    
        return cls(router=router, experts=experts, expert_names=names, k=cfg.get("k", 2), tau=cfg.get("tau", 1.0), device=device)

    def forward(self, x, k=None):
        if k is None:
            k = self.k
        router_logits = self.router(x)
        topk_scores, topk_indices = torch.topk(router_logits, k=k, dim=-1)
        gate_weights = F.softmax(topk_scores / self.tau, dim=-1)
        return router_logits, topk_indices, gate_weights