KridgeDookie's picture
Add validated original-size base plus compact rank-8 correction
6f27bbf verified
Raw History Blame Contribute Delete
4.79 kB
"""Approximate the verified mixed-Q8 edit as a small additive adapter.
This is a new candidate, not an exact reconstruction of the BF16 surgery.
Only exported-model behavioral evaluation may qualify it for publication.
"""
import json, math, time
from pathlib import Path
import numpy as np
import torch
from safetensors import safe_open
from safetensors.torch import save_file
from gguf import GGUFWriter
R=Path('/workspace/bonsai2')
if (R/'ADAPTERS_READY').exists():raise SystemExit(0)
torch.manual_seed(1337)
torch.set_num_threads(8)
torch.backends.cuda.matmul.allow_tf32=False
cfg=json.loads((R/'source_mlx/config.json').read_text())
src=safe_open(R/'source_mlx/model.safetensors',framework='pt',device='cpu')
ref=safe_open(R/'reference/model.safetensors',framework='pt',device='cpu')
ranks=(8,16,32)
packs={rank:{} for rank in ranks}
writers={}
out=R/'adapters';out.mkdir(exist_ok=True)
for rank in ranks:
w=GGUFWriter(str(out/f'bonsai2-compact-r{rank}.gguf'),'qwen35')
w.add_type('adapter');w.add_string('adapter.type','lora')
w.add_float32('adapter.lora.alpha',float(rank))
w.add_name(f'Bonsai 2 Philadelphia compact rank {rank} correction')
writers[rank]=w
def decode(f,key,bits,group):
p=f.get_tensor(key+'.weight').to(device='cuda',dtype=torch.int64)
s=f.get_tensor(key+'.scales').cuda().float()
b=f.get_tensor(key+'.biases').cuda().float()
shifts=bits*torch.arange(32//bits,device='cuda',dtype=torch.int64)
x=((p[...,None]>>shifts)&((1<<bits)-1)).float().reshape(p.shape[0],-1,group)
return (x*s[...,None]+b[...,None]).reshape(p.shape[0],-1)
def inverse_hadamard(x,block,signs):
shape=x.shape;z=x.reshape(-1,block).clone();h=1
while h<block:
v=z.reshape(-1,block//(2*h),2*h)
a,b=v[...,:h].clone(),v[...,h:].clone()
v[...,:h]=a+b;v[...,h:]=a-b;h*=2
return z.reshape(shape)/math.sqrt(block)*signs
rows=[];started=time.time()
for record in cfg['modules']:
path=record['path'];parts=path.split('.')
if len(parts)!=5 or parts[:2]!=['model','layers']:continue
layer=int(parts[2]);component='.'.join(parts[3:])
stem={'mlp.down_proj':'ffn_down','self_attn.o_proj':'attn_output','linear_attn.out_proj':'ssm_out'}.get(component)
if layer==0 or stem is None:continue
key='language_model.'+path
original=decode(src,key,2,128)
if record['block']:
original=inverse_hadamard(original,record['block'],src.get_tensor(key+'.signs').cuda().float())
target=decode(ref,key,8,32)
delta=target-original
u,s,v=torch.svd_lowrank(delta,q=40,niter=4)
order=torch.argsort(s,descending=True);u=u[:,order];s=s[order];v=v[:,order]
entry={'module':path,'shape':list(delta.shape),'delta_relative_norm':float(delta.norm()/target.norm()),'ranks':{}}
probe=torch.randn(16,delta.shape[1],device='cuda')
for rank in ranks:
root=s[:rank].sqrt()
a=(v[:,:rank]*root).T.contiguous().half()
b=(u[:,:rank]*root).contiguous().half()
assert torch.isfinite(a).all() and torch.isfinite(b).all()
approx=b.float()@a.float()
entry['ranks'][str(rank)]={'residual_relative_to_target':float((delta-approx).norm()/target.norm()),'delta_energy_captured':float(1-(delta-approx).square().sum()/delta.square().sum()),'probe_relative_error':float(((probe@original.T+(probe@a.float().T)@b.float().T)-(probe@target.T)).norm()/(probe@target.T).norm())}
packs[rank][key+'.lora_a']=a.cpu();packs[rank][key+'.lora_b']=b.cpu()
native_a=a
if component=='linear_attn.out_proj':
perm=torch.arange(a.shape[1],device='cuda').reshape(3,16,128).transpose(0,1).reshape(-1)
native_a=a[:,torch.argsort(perm)]
name=f'blk.{layer}.{stem}.weight'
writers[rank].add_tensor(name+'.lora_a',native_a.cpu().numpy())
writers[rank].add_tensor(name+'.lora_b',b.cpu().numpy())
rows.append(entry)
print('COMPRESSED',len(rows),path,entry['ranks']['8'],flush=True)
del original,target,delta,u,s,v,approx
assert len(rows)==126,len(rows)
for rank in ranks:
save_file(packs[rank],str(out/f'bonsai2-compact-r{rank}.safetensors'))
w=writers[rank];w.write_header_to_file();w.write_kv_data_to_file();w.write_tensors_to_file();w.close()
receipt={'method':'randomized SVD of canonical mixed-Q8 minus decoded original ternary weights','seed':1337,'svd_q':40,'svd_power_iterations':4,'ranks':ranks,'modules':126,'source_revision':'3f926b415992eaa2ae9dd7b573706494d6bbf787','reference_revision':'61231c9a2feed6457f47d1d50a691b8dc96b2613','gguf_adapter_input_basis':'canonical, cyclic GDN columns','mlx_adapter_input_basis':'canonical, grouped GDN columns','seconds':time.time()-started,'rows':rows}
(out/'compression.json').write_text(json.dumps(receipt,indent=2))
(R/'ADAPTERS_READY').touch()