File size: 2,867 Bytes
6f27bbf
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""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))