Ling-3.0-tiny-RKNN / tools /audit_absolute_weights.py
Sariel00's picture
Publish Ling-3.0-tiny RKNN engine and model
3fd1a35 verified
Raw History Blame Contribute Delete
8.12 kB
#!/usr/bin/env python3
"""Audit actual exported linear weights against pinned BF16, without inference.
Read in row chunks; exclude alignment padding, raw embeddings, routers and norms.
FP64 reductions; W8 reconstruction follows src/linear.cpp (FP32 per-row scale).
No activation quantization or kernel errors are represented by these metrics.
"""
import argparse
import json
import math
import mmap
import time
from pathlib import Path
import numpy as np
import torch
from pack_model import HEADER, ENTRY
from quantize_model import TensorSource, source_linear_names
class Package:
def __init__(self, path):
self.file = path.open('rb')
self.data = mmap.mmap(self.file.fileno(), 0, access=mmap.ACCESS_READ)
h = HEADER.unpack_from(self.data)
assert h[0] == b'L3RKNN1\0' and h[12] == len(self.data)
self.info = dict(path=str(path.resolve()), bytes=len(self.data), revision=h[13].hex())
self.entries = {}
for i in range(h[4]):
e = ENTRY.unpack_from(self.data, h[7] + i * ENTRY.size)
name = self.data[h[8]+e[0]:h[8]+e[0]+e[1]].decode()
self.entries[name] = e
def array(self, name, dtype):
e = self.entries[name]
dt = np.dtype(dtype)
return np.frombuffer(self.data, dtype=dt, count=e[13]//dt.itemsize, offset=e[12])
def rows(self, base, lo, hi):
e = self.entries[base+'.weight']
k, n = e[8:10]
if e[2] == 1:
assert e[6] == 0
b = self.array(base+'.weight', '<u2').reshape(n, k)[lo:hi]
return (b.astype(np.uint32) << 16).view(np.float32)
assert e[2] == 5 and e[6] == 2 and lo % 2 == 0 and hi % 2 == 0
b = self.array(base+'.weight', 'u1').reshape(k, n//2)[:, lo//2:hi//2]
q = np.empty((k, hi-lo), dtype=np.int8)
q[:, 0::2] = b & 15
q[:, 1::2] = b >> 4
q[q >= 8] -= 16
q = q.T.astype(np.float32)
if e[5] == 4:
assert e[7] == 32
b = self.array(base+'.scales', '<u2').reshape(n, k//32)[lo:hi]
s = (b.astype(np.uint32) << 16).view(np.float32)
return (q.reshape(hi-lo, k//32, 32) * s[:, :, None]).reshape(hi-lo, k)
assert e[5] == 2
return q * self.array(base+'.scales', '<f4')[lo:hi, None]
def family(base):
if base == 'lm_head': return 'head'
if '.attention.' in base:
kind = 'mla' if (int(base.split('.')[2])+1) % 4 == 0 else 'kda'
return kind + ('_out' if base.endswith('.o_proj') else '_in')
if '.shared_experts.' in base: return 'shared'
if '.mlp.experts.' in base: return 'routed'
if base.startswith('model.layers.0.mlp.'): return 'dense'
raise ValueError(base)
def w8(w):
s = np.max(np.abs(w), axis=1) / np.float32(127)
s[s == 0] = 1
return np.clip(np.rint(w / s[:, None]), -127, 127) * s[:, None]
def blank():
return dict(count=0, weight_sum_sq=0., error_sum_sq=0., error_sum_abs=0., max_abs_error=0.)
def add(a, b):
for key in a:
a[key] = max(a[key], b[key]) if key == 'max_abs_error' else a[key]+b[key]
def measure(ref, candidate):
# Subtraction in float64 also preserves differences across tiny/large values.
r = ref.astype(np.float64)
d = candidate.astype(np.float64)-r
assert np.isfinite(r).all() and np.isfinite(d).all()
return dict(count=r.size, weight_sum_sq=float(np.square(r).sum()),
error_sum_sq=float(np.square(d).sum()), error_sum_abs=float(np.abs(d).sum()),
max_abs_error=float(np.abs(d).max()))
def finish(a):
return dict(**a, mae=a['error_sum_abs']/a['count'], mse=a['error_sum_sq']/a['count'],
rmse=math.sqrt(a['error_sum_sq']/a['count']),
relative_frobenius_rmse=math.sqrt(a['error_sum_sq']/a['weight_sum_sq']),
snr_db=10*math.log10(a['weight_sum_sq']/a['error_sum_sq']) if a['error_sum_sq'] else None)
def main():
p = argparse.ArgumentParser(description=__doc__)
p.add_argument('--source', type=Path, required=True)
p.add_argument('--original', type=Path, required=True)
p.add_argument('--calibrated', type=Path, required=True)
p.add_argument('--official', type=Path, required=True)
p.add_argument('--output', type=Path, required=True)
a = p.parse_args()
torch.set_num_threads(2)
src = TensorSource(a.source)
packs = {name: Package(getattr(a, name)) for name in ('original', 'calibrated', 'official')}
policies = dict(original_w4=set(), scheme3={'kda_in','kda_out','mla_in','mla_out','shared'},
scheme4={'kda_in','kda_out','mla_in','mla_out','shared','dense','head'},
selective={'kda_out','mla_in','mla_out','shared','dense','head'})
variants = list(policies) + ['official_int4_exact', 'official_w8_bridge']
totals = {v: {} for v in variants}
matrices = []
started = time.monotonic()
for name, entry in packs['original'].entries.items():
if entry[4] != 3: continue
base = name.removesuffix('.weight')
f = family(base)
for pack in packs.values(): assert pack.entries[name][8:10] == entry[8:10]
k, padded_n = entry[8:10]
acc = {key: blank() for key in ('original', 'official', 'bridge', 'calibrated')}
offset = 0
for key in source_linear_names(base):
tensor = src.tensor(key)
assert tensor.dtype == torch.bfloat16 and tensor.shape[1] == k and tensor.shape[0] % 2 == 0
for lo in range(0, tensor.shape[0], 256):
hi = min(lo+256, tensor.shape[0])
ref = tensor[lo:hi].float().numpy()
original = packs['original'].rows(base, offset+lo, offset+hi)
official = packs['official'].rows(base, offset+lo, offset+hi)
if f != 'routed': assert np.array_equal(ref, official), (base, lo)
add(acc['original'], measure(ref, original))
add(acc['official'], measure(ref, official))
add(acc['bridge'], measure(ref, w8(official)))
if f == 'kda_in':
add(acc['calibrated'], measure(ref, packs['calibrated'].rows(base, offset+lo, offset+hi)))
offset += tensor.shape[0]
assert offset <= padded_n
selected = {}
for v in variants:
if v == 'official_int4_exact': source = 'official'
elif v == 'official_w8_bridge' or f in policies.get(v, set()): source = 'bridge'
elif v == 'selective' and f == 'kda_in': source = 'calibrated'
else: source = 'original'
selected[v] = source
add(totals[v].setdefault(f, blank()), acc[source])
matrices.append(dict(name=base, family=f, layer=None if base == 'lm_head' else int(base.split('.')[2]),
rows=offset, columns=k, padded_rows=padded_n, selected=selected,
metrics={key:finish(value) for key,value in acc.items() if value['count']}))
if len(matrices) % 256 == 0:
print(json.dumps(dict(matrices=len(matrices), elapsed_s=time.monotonic()-started)), flush=True)
result = dict(scope='All exported linear source elements, padding excluded; not whole-model capability loss',
source=str(a.source.resolve()), packages={key:pack.info for key,pack in packs.items()},
elapsed_s=time.monotonic()-started, matrix_count=len(matrices), variants={}, matrices=matrices)
for v, families in totals.items():
total, nonrouted = blank(), blank()
for f, value in families.items():
add(total, value)
if f != 'routed': add(nonrouted, value)
result['variants'][v] = dict(all_linears=finish(total), nonrouted=finish(nonrouted),
families={f:finish(value) for f,value in families.items()})
a.output.parent.mkdir(parents=True, exist_ok=True)
a.output.write_text(json.dumps(result, indent=2)+'\n')
print(json.dumps({v:x['all_linears'] for v,x in result['variants'].items()}), flush=True)
if __name__ == '__main__': main()