Sparse-AST-BWM / model.py
Chaman1234's picture
Publish unified Sparse-AST / BWM model family under one repository
fec7158 verified
Raw
History Blame Contribute Delete
7.56 kB
"""
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