Download variants/mlx-decoder/decoder.py from phequals/indic-transcribe-flex-coreml: direct link, hf CLI and curl.
- Browser
- Download file 4.5 kB
-
https://huggingface.co/phequals/indic-transcribe-flex-coreml/resolve/main/variants/mlx-decoder/decoder.py
- Command line
-
hf download hf://phequals/indic-transcribe-flex-coreml/variants/mlx-decoder/decoder.py
-
curl -L -o decoder.py https://huggingface.co/phequals/indic-transcribe-flex-coreml/resolve/main/variants/mlx-decoder/decoder.py
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) | |