phequals's picture
Add validated CoreML encoder and MLX decoder prototype
ca1e261 verified
Raw
History Blame Contribute Delete
4.5 kB
"""Bodhan decoder-only MLX prototype. Original checkpoint naming and scaling."""
from pathlib import Path
import json, math, time
import mlx.core as mx
class Decoder:
def __init__(self, root):
root=Path(root)
self.config=json.loads((root/'config.json').read_text())
self.w=mx.load(str(root/'decoder.safetensors'))
mx.eval(self.w)
self.heads=self.config['decoder_attention_heads']
self.dim=self.config['d_model']//self.heads
self.layers=self.config['decoder_layers']
self.attn_scale=math.sqrt(math.sqrt(self.dim))
def linear(self,x,prefix):
weight=self.w[prefix+'.weight'];bias=self.w.get(prefix+'.bias')
return x @ weight.T if bias is None else mx.addmm(bias,x,weight.T)
def norm(self,x,prefix):
return mx.fast.layer_norm(x,self.w[prefix+'.weight'],self.w[prefix+'.bias'],1e-5)
def split(self,x):
return x.reshape(x.shape[0],x.shape[1],self.heads,self.dim).transpose(0,2,1,3)
def project(self,x,prefix):
return (self.split(self.linear(x,prefix+'.key_net'))/self.attn_scale,
self.split(self.linear(x,prefix+'.value_net')))
def attention(self,x,kv,prefix,mask=None):
q=self.split(self.linear(x,prefix+'.query_net'))/self.attn_scale
output=mx.fast.scaled_dot_product_attention(q,*kv,scale=1.0,mask=mask)
output=output.transpose(0,2,1,3).reshape(x.shape[0],x.shape[1],-1)
return self.linear(output,prefix+'.out_projection')
def cross_cache(self,encoder):
result=[self.project(encoder,f'layers.{i}.second_sub_layer') for i in range(self.layers)]
mx.eval(result)
return result
def step(self,ids,position,cross,cache):
positions=mx.arange(position,position+ids.shape[1])
x=self.w['embedding.token_embedding.weight'][ids]+self.w['embedding.position_embedding.pos_enc'][positions][None]
x=self.norm(x,'embedding.layer_norm')
# Only populated cache positions participate in attention.
mask='causal' if ids.shape[1]>1 else None
new_cache=[]
for i in range(self.layers):
p=f'layers.{i}'
n=self.norm(x,p+'.layer_norm_1')
k,v=self.project(n,p+'.first_sub_layer')
if cache is not None:
k=mx.concatenate([cache[i][0],k],axis=2)
v=mx.concatenate([cache[i][1],v],axis=2)
new_cache.append((k,v))
x=x+self.attention(n,(k,v),p+'.first_sub_layer',mask)
x=x+self.attention(self.norm(x,p+'.layer_norm_2'),cross[i],p+'.second_sub_layer')
x=x+self.linear(mx.maximum(self.linear(self.norm(x,p+'.layer_norm_3'),p+'.third_sub_layer.dense_in'),0),p+'.third_sub_layer.dense_out')
return self.linear(self.norm(x,'final_layer_norm'),'lm_head'),new_cache
def generate(self,encoder,tokenizer,language='hi',mixed=False,max_tokens=256):
start=time.perf_counter();cross=self.cross_cache(encoder);cross_seconds=time.perf_counter()-start
start=time.perf_counter()
if language=='auto':
prefix=tokenizer['prompts']['hi'][:3]
logits,_=self.step(mx.array([prefix],dtype=mx.int32),0,cross,None)
scores=logits[0,-1];mx.eval(scores)
language=max(tokenizer['prompts'],key=lambda k:float(scores[tokenizer['prompts'][k][3]].item()))
prompts=tokenizer['mixed_prompts'] if mixed else tokenizer['prompts']
ids=mx.array([prompts[language]],dtype=mx.int32)
position=0;cache=None;tokens=[];prefill_seconds=None
for index in range(max_tokens):
logits,cache=self.step(ids,position,cross,cache)
next_token=mx.argmax(logits[0,-1],axis=-1)
mx.eval(next_token,cache)
next_token=int(next_token.item())
if index==0:prefill_seconds=time.perf_counter()-start
position+=ids.shape[1]
if next_token==tokenizer['eos_id']:break
tokens.append(next_token);ids=mx.array([[next_token]],dtype=mx.int32)
else:raise RuntimeError('No EOS within token limit; refusing truncated transcript')
seconds=time.perf_counter()-start
text=''.join(tokenizer['pieces'][t] for t in tokens if t>=tokenizer['special_count']).replace('▁',' ').strip()
return dict(text=text,language=language,tokens=len(tokens),endedWithEOS=True,crossSeconds=cross_seconds,decodeSeconds=seconds,prefillSeconds=prefill_seconds,generationSeconds=seconds-prefill_seconds)