File size: 11,100 Bytes
7aa9bc4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Exact 2,990,000-parameter Sol Milkshake model implemented in MLX."""
from __future__ import annotations
import math
from collections import OrderedDict
import mlx.core as mx
import mlx.nn as nn
from mlx.utils import tree_flatten
from sol_config import DEPLOYED_PARAMS, SolConfig

def unit_norm(x):
    xf=x.astype(mx.float32)
    return (xf*mx.rsqrt(mx.maximum(mx.sum(xf*xf,axis=-1,keepdims=True),mx.array(1e-12)))).astype(x.dtype)

def _count(module): return sum(int(x.size) for _,x in tree_flatten(module.trainable_parameters()))

def _rope(q,k,theta):
    n,d=q.shape[-2:]; inv=mx.exp(mx.arange(0,d,2,dtype=mx.float32)*(-math.log(theta)/d))
    ang=mx.arange(n,dtype=mx.float32)[:,None]*inv[None,:]; cs,sn=mx.cos(ang)[None,None],mx.sin(ang)[None,None]
    def rot(x):
        xf=x.astype(mx.float32); a,b=xf[...,::2],xf[...,1::2]
        return mx.stack((a*cs-b*sn,a*sn+b*cs),axis=-1).reshape(x.shape).astype(x.dtype)
    return rot(q),rot(k)

class TiedTokens(nn.Module):
    def __init__(self,c):
        super().__init__(); self.weight=unit_norm(mx.random.normal((c.vocab_size,c.d_model)))
    def encode(self,ids): return self.weight[ids]
    def decode(self,x,scale): return x.astype(mx.float32)@unit_norm(self.weight).astype(mx.float32).T*mx.exp(scale)

class TNGram(nn.Module):
    """Causal hashed CP retrieval for orders 2--5 (exactly 573,460 parameters)."""
    def __init__(self,c):
        super().__init__(); r=c.tn_rank; self.token_factor=mx.random.normal((c.vocab_size,r))*.02
        self.hash_tables=mx.random.normal((4,4950,r))*.02; self.order_factors=mx.ones((4,r))
        self.rank_to_hidden=mx.random.normal((r,c.d_model))*.02; self.output_bias=mx.zeros((c.d_model,))
        self.gates=mx.full((4,),-2.0); self.order_mix=mx.eye(4,r); self.context_controls=mx.zeros((16,))
    def __call__(self,ids):
        b,t=ids.shape; out=mx.broadcast_to(self.output_bias,(b,t,self.output_bias.shape[0]))
        for oi,order in enumerate(range(2,6)):
            h=mx.zeros_like(ids); prod=mx.ones((b,t,self.token_factor.shape[1]),dtype=mx.float32); valid=mx.arange(t)>=order-1
            for lag in range(1,order):
                shifted=mx.zeros((b,t),dtype=ids.dtype) if lag>=t else mx.concatenate((mx.zeros((b,lag),dtype=ids.dtype),ids[:,:-lag]),axis=1)
                h=(h*(131+oi*6)+shifted)%4950; prod=prod*self.token_factor[shifted].astype(mx.float32)
            rank=prod*self.hash_tables[oi,h]*self.order_factors[oi]
            rank=rank+self.order_mix[oi][None,None,:]*mx.mean(mx.tanh(self.context_controls[oi*4:(oi+1)*4]))
            out=out+(rank@self.rank_to_hidden)*mx.sigmoid(self.gates[oi])*valid[None,:,None]
        return out

class ShakeMemory(nn.Module):
    """Completed-chunk rolling memory (exactly 51,520 parameters)."""
    def __init__(self,c):
        super().__init__(); w=c.memory_width; self.chunk_size=c.chunk_size; self.slots=c.memory_slots
        self.summary_router=nn.Linear(c.d_model,1,bias=False); self.compress=nn.Linear(c.d_model,w,bias=False)
        self.query=nn.Linear(c.d_model,w,bias=False); self.key=nn.Linear(w,w,bias=False)
        self.value=nn.Linear(w,w,bias=False); self.output=nn.Linear(w,c.d_model,bias=False)
        self.slot_positions=mx.random.normal((c.memory_slots,c.d_model))*.01
        self.decay=mx.zeros((w,)); self.update_gate=mx.zeros((w,))
    def __call__(self,x):
        b,t,d=x.shape; chunks=[]; reads=[]
        for start in range(0,t,self.chunk_size):
            end=min(t,start+self.chunk_size); length=end-start
            if chunks:
                mem=mx.stack(chunks[-self.slots:],axis=1); count=mem.shape[1]
                mem=mem+self.compress(self.slot_positions[:count])[None]*mx.sigmoid(self.decay)[None,None]
                a=mx.softmax((self.query(x[:,start:end]).astype(mx.float32)@self.key(mem).astype(mx.float32).transpose(0,2,1))/math.sqrt(mem.shape[-1]),axis=-1)
                reads.append(self.output((a@self.value(mem).astype(mx.float32)).astype(x.dtype)))
            else: reads.append(mx.zeros((b,length,d),dtype=x.dtype))
            if length==self.chunk_size:
                weights=mx.softmax(self.summary_router(x[:,start:end]).astype(mx.float32).squeeze(-1),axis=-1)
                summary=mx.sum(weights[:,:,None]*x[:,start:end].astype(mx.float32),axis=1)
                chunks.append(self.compress(summary.astype(x.dtype))*mx.sigmoid(self.update_gate))
        return mx.concatenate(reads,axis=1)

class Attention(nn.Module):
    def __init__(self,c):
        super().__init__(); d=c.d_model; self.q=nn.Linear(d,d,bias=False); self.k=nn.Linear(d,c.n_kv_heads*c.head_dim,bias=False)
        self.v=nn.Linear(d,c.n_kv_heads*c.head_dim,bias=False); self.o=nn.Linear(d,d,bias=False)
        self.hq=c.n_q_heads; self.hk=c.n_kv_heads; self.hd=c.head_dim; self.theta=c.rope_theta
        self.value_current=mx.zeros((1,)); self.value_prior=mx.zeros((1,))
    def __call__(self,x,prior=None):
        b,t,_=x.shape; q=unit_norm(self.q(x).reshape(b,t,self.hq,self.hd)).transpose(0,2,1,3)
        k=unit_norm(self.k(x).reshape(b,t,self.hk,self.hd)).transpose(0,2,1,3); v=self.v(x).reshape(b,t,self.hk,self.hd).transpose(0,2,1,3)
        if prior is not None: v=mx.sigmoid(self.value_current)*v+mx.sigmoid(self.value_prior)*prior
        q,k=_rope(q,k,self.theta); k=mx.repeat(k,self.hq//self.hk,axis=1); vr=mx.repeat(v,self.hq//self.hk,axis=1)
        scores=(q.astype(mx.float32)@k.astype(mx.float32).transpose(0,1,3,2))/math.sqrt(self.hd)
        pos=mx.arange(t); mask=mx.where(pos[:,None]>=pos[None,:],0.0,-1e9)
        attended=(mx.softmax(scores+mask[None,None],axis=-1)@vr.astype(mx.float32)).astype(x.dtype)
        vh=unit_norm(vr); attended=attended-mx.sum(attended.astype(mx.float32)*vh.astype(mx.float32),axis=-1,keepdims=True).astype(x.dtype)*vh
        return self.o(attended.transpose(0,2,1,3).reshape(b,t,-1)),v

class Block(nn.Module):
    def __init__(self,c):
        super().__init__(); self.attention=Attention(c); self.gate=nn.Linear(c.d_model,c.ffn_hidden,bias=False)
        self.up=nn.Linear(c.d_model,c.ffn_hidden,bias=False); self.down=nn.Linear(c.ffn_hidden,c.d_model,bias=False)
        self.attn_step=mx.zeros((c.d_model,)); self.ffn_step=mx.zeros((c.d_model,))
    def __call__(self,x,prior=None):
        a,v=self.attention(unit_norm(x),prior); x=unit_norm(x+mx.sigmoid(self.attn_step)*(unit_norm(a)-x))
        f=self.down(nn.silu(self.gate(x))*self.up(x)); return unit_norm(x+mx.sigmoid(self.ffn_step)*(unit_norm(f)-x)),v

class SolMilkshake(nn.Module):
    def __init__(self,c=None):
        super().__init__(); self.config=c or SolConfig(); self.tokens=TiedTokens(self.config); self.tn_gram=TNGram(self.config)
        self.shake_memory=ShakeMemory(self.config); self.blocks=[Block(self.config) for _ in range(5)]
        self.pass_embeddings=mx.zeros((3,self.config.d_model)); self.mod_routers=[nn.Linear(self.config.d_model,1) for _ in range(6)]
        self.logit_scale=mx.zeros((self.config.vocab_size,)); self.misc_controls=mx.zeros((12,)); self._verify_ledger()
    def _select(self,x,old,router,capacity,lam):
        delta=1-mx.sum(unit_norm(x).astype(mx.float32)*unit_norm(old).astype(mx.float32),axis=-1)
        score=router(x).squeeze(-1)+lam*delta
        positions=mx.arange(x.shape[1])
        mask=(positions%4!=3) if capacity==.75 else (positions%2==0)
        return mx.broadcast_to(mask[None,:],score.shape),mx.sigmoid(score)
    def __call__(self,ids):
        x=self.tokens.encode(ids); x=unit_norm(x+mx.sigmoid(self.misc_controls[0])*self.tn_gram(ids))
        x=unit_norm(x+mx.sigmoid(self.misc_controls[1])*self.shake_memory(x)); x,v=self.blocks[0](x); ri=0
        for p in range(3):
            old=x; x=unit_norm(x+self.pass_embeddings[p]); prior=None
            for bi in range(1,4):
                proposal,prior=self.blocks[bi](x,prior)
                if p==0: x=proposal
                else:
                    mask,route_gate=self._select(x,old,self.mod_routers[ri],.75 if p==1 else .5,self.misc_controls[2+ri]); ri+=1
                    routed=unit_norm(x+route_gate[:,:,None]*(proposal-x))
                    x=mx.where(mask[:,:,None],routed,x)
        x,_=self.blocks[4](x,v); return self.tokens.decode(x,self.logit_scale)
    def parameter_ledger(self):
        ledger=OrderedDict(token_embedding_tied_head=int(self.tokens.weight.size),attention_total=sum(_count(b.attention)-2 for b in self.blocks),
          swiglu_total=sum(_count(b.gate)+_count(b.up)+_count(b.down) for b in self.blocks),tn_gram=_count(self.tn_gram),shake_memory=_count(self.shake_memory),
          ngpt_steps=sum(int(b.attn_step.size+b.ffn_step.size) for b in self.blocks),pass_embeddings=int(self.pass_embeddings.size),
          mod_routers=sum(_count(r) for r in self.mod_routers),logit_scale=int(self.logit_scale.size),value_residual=10,
          misc_controls=int(self.misc_controls.size),xsa=0); ledger['deployed_total']=sum(ledger.values()); return dict(ledger)
    def deployed_parameter_count(self): return _count(self)
    def _verify_ledger(self):
        expected={"token_embedding_tied_head":393216,"attention_total":491520,"swiglu_total":1474560,"tn_gram":573460,"shake_memory":51520,"ngpt_steps":1920,"pass_embeddings":576,"mod_routers":1158,"logit_scale":2048,"value_residual":10,"misc_controls":12,"xsa":0,"deployed_total":DEPLOYED_PARAMS}
        if self.parameter_ledger()!=expected or self.deployed_parameter_count()!=DEPLOYED_PARAMS: raise RuntimeError(f"parameter ledger mismatch: {self.parameter_ledger()}, actual={self.deployed_parameter_count()}")

def future_token_causality_error(model,sequence_length=35):
    ids=mx.arange(sequence_length,dtype=mx.int32)[None]%model.config.vocab_size; boundary=sequence_length//2
    changed=mx.concatenate((ids[:,:boundary],(ids[:,boundary:]+17)%model.config.vocab_size),axis=1)
    a,b=model(ids),model(changed); mx.eval(a,b); return float(mx.max(mx.abs(a[:,:boundary]-b[:,:boundary])))

def load_model(model_dir):
    """Load the standalone MLX model and frozen tokenizer from a snapshot."""
    from pathlib import Path
    from tokenizers import Tokenizer

    model_dir = Path(model_dir)
    model = SolMilkshake(SolConfig())
    model.load_weights(str(model_dir / "model.npz"))
    mx.eval(model.parameters())
    tokenizer = Tokenizer.from_file(str(model_dir / "tokenizer.json"))
    return model, tokenizer

def generate(model, tokenizer, prompt, max_new_tokens=64, temperature=0.0, seed=0):
    """Generate a single continuation with greedy or temperature sampling."""
    mx.random.seed(seed)
    tokens = tokenizer.encode(prompt).ids
    if not tokens:
        tokens = [0]
    for _ in range(max_new_tokens):
        context = tokens[-model.config.max_context:]
        logits = model(mx.array([context], dtype=mx.int32))[0, -1]
        if temperature and temperature > 0:
            next_token = int(mx.random.categorical(logits / temperature).item())
        else:
            next_token = int(mx.argmax(logits).item())
        tokens.append(next_token)
        if next_token == 1:
            break
    return tokenizer.decode(tokens)