KridgeDookie's picture
Add validated original-size base plus compact rank-8 correction
6f27bbf verified
Raw History Blame Contribute Delete
2.87 kB
"""Load the original Bonsai MLX pack with the required compact correction.
The adapter remains separate: merging it and requantizing destroys its edits.
"""
import sys
from pathlib import Path
import mlx.core as mx
from mlx import nn
class CorrectedLinear(nn.Module):
def __init__(self,base,a,b):
super().__init__()
self.base,self.a,self.b=base,a,b
def __call__(self,x):
y=self.base(x)
correction=(x.astype(mx.float32)@self.a.astype(mx.float32).T)@self.b.astype(mx.float32).T
return (y.astype(mx.float32)+correction).astype(y.dtype)
def load_compact(directory,adapter=None,load_processor=True):
directory=Path(directory).resolve()
sys.path.insert(0,str(directory/'runtime'))
from vision_artifact import load_vl_model
model,processor,config=load_vl_model(directory,load_processor=load_processor)
adapter=Path(adapter) if adapter else directory/'adapter.safetensors'
weights=mx.load(str(adapter))
paths={k.removesuffix('.lora_a') for k in weights if k.endswith('.lora_a')}
if len(paths)!=126 or len(weights)!=252:
raise ValueError('Expected all 126 correction pairs')
for path in sorted(paths):
parts=path.split('.');parent=model
for part in parts[:-1]:parent=parent[int(part)] if part.isdigit() else getattr(parent,part)
base=getattr(parent,parts[-1]);a,b=weights[path+'.lora_a'],weights[path+'.lora_b']
if a.ndim!=2 or b.ndim!=2 or a.shape[0]!=b.shape[1]:raise ValueError('Invalid correction shape: '+path)
if base.weight.shape[0]!=b.shape[0] or base.weight.shape[1]*16!=a.shape[1]:raise ValueError('Wrong base pack: '+path)
setattr(parent,parts[-1],CorrectedLinear(base,a,b))
mx.eval(model.parameters());model.eval()
return model,processor,config
if __name__=='__main__':
import argparse
from transformers import AutoTokenizer
p=argparse.ArgumentParser(description=__doc__)
p.add_argument('--model',default=str(Path(__file__).resolve().parent))
p.add_argument('--adapter')
p.add_argument('--prompt',required=True)
p.add_argument('--max-tokens',type=int,default=256)
args=p.parse_args()
model,_,_=load_compact(args.model,args.adapter,load_processor=False)
tokenizer=AutoTokenizer.from_pretrained(args.model)
text=tokenizer.apply_chat_template([{'role':'user','content':args.prompt}],tokenize=False,add_generation_prompt=True,enable_thinking=False)
x=mx.array([tokenizer.encode(text,add_special_tokens=False)])
lm=model.language_model;cache=lm.make_cache();tokens=[]
for _ in range(args.max_tokens):
logits=lm(x,cache=cache).logits[:,-1,:]
token=int(mx.argmax(logits,axis=-1).item())
if token==tokenizer.eos_token_id:break
tokens.append(token);x=mx.array([[token]])
print(tokenizer.decode(tokens,skip_special_tokens=True))