"""Synthetic n-gram LM: symmetric RTN, real packing, FP32 dequantized compute.
No GPTQ/AWQ algorithm, GPU memory measurement, or low-bit compute kernel.
"""
import os
for name in ('OMP_NUM_THREADS', 'OPENBLAS_NUM_THREADS', 'VECLIB_MAXIMUM_THREADS'):
    os.environ.setdefault(name, '1')
import argparse
import hashlib
import json
import platform
import sys
import time
from pathlib import Path
import numpy as np

V, D, H = 16, 8, 64
WEIGHTS = ('E', 'W1', 'W2')


def data(seed, n):
    rng = np.random.default_rng(seed)
    seq = rng.integers(V, size=(n, 34))
    for t in range(2, 34):
        rule = (seq[:, t-1] + 1 + seq[:, t-2] % 2) % V
        seq[:, t] = np.where(rng.random(n) < .9, rule, rng.integers(V, size=n))
    x = np.stack((seq[:, :-2], seq[:, 1:-1]), axis=-1).reshape(-1, 2)
    return seq, x, seq[:, 2:].reshape(-1)


def initialize(seed):
    rng = np.random.default_rng(seed)
    return dict(E=rng.normal(0, .2, (V, D)).astype('float32'),
                W1=rng.normal(0, .2, (2*D, H)).astype('float32'),
                W2=rng.normal(0, .2, (H, V)).astype('float32'),
                b1=np.zeros(H, 'float32'), b2=np.zeros(V, 'float32'))


def forward(p, x):
    z = p['E'][x].reshape(len(x), 2*D)
    h = np.tanh(z @ p['W1'] + p['b1'])
    logits = h @ p['W2'] + p['b2']
    return logits, z, h


def loss_grad(p, x, y):
    logits, z, h = forward(p, x)
    shift = logits - logits.max(axis=1, keepdims=True)
    logp = shift - np.log(np.exp(shift).sum(axis=1, keepdims=True))
    loss = -logp[np.arange(len(y)), y].mean()
    dl = np.exp(logp)
    dl[np.arange(len(y)), y] -= 1
    dl /= len(y)
    dh = (dl @ p['W2'].T) * (1-h*h)
    dz = (dh @ p['W1'].T).reshape(len(x), 2, D)
    de = np.zeros_like(p['E'])
    np.add.at(de, x.reshape(-1), dz.reshape(-1, D))
    return float(loss), dict(E=de, W1=z.T@dh, W2=h.T@dl,
                            b1=dh.sum(0), b2=dl.sum(0))


def nlls(p, x, y):
    a = forward(p, x)[0].astype('float64')
    a -= a.max(axis=1, keepdims=True)
    return np.log(np.exp(a).sum(1)) - a[np.arange(len(y)), y]


def pack(w, bits, group):
    assert bits in (4, 8) and w.size % group == 0
    rows = w.reshape(-1, group)
    maxq = 2**(bits-1)-1
    scales = (np.abs(rows).max(axis=1)/maxq).astype('<f4')
    scales[scales == 0] = 1
    q = np.clip(np.rint(rows/scales[:, None]), -maxq, maxq).astype('int8').ravel()
    if bits == 4:
        u = (q.astype('int16') + 8).astype('uint8')
        payload = (u[::2] | (u[1::2] << 4)).tobytes()
    else:
        payload = q.tobytes()
    return dict(bits=bits, group=group, shape=list(w.shape), payload=payload,
                scales=scales, q=q)


def unpack(a):
    if a['bits'] == 4:
        b = np.frombuffer(a['payload'], dtype='uint8')
        q = np.stack((b & 15, b >> 4), axis=1).reshape(-1).astype('int16')-8
    else:
        q = np.frombuffer(a['payload'], dtype='int8')
    w = q.astype('float32').reshape(-1, a['group']) * a['scales'][:, None]
    return w.reshape(a['shape'])


def restore(q, p):
    return {k:unpack(q[k]) if k in q else p[k] for k in p}


def benchmark(fn, repeats=101):
    for _ in range(20): fn()
    measurements = []
    for _ in range(repeats):
        start = time.perf_counter_ns()
        fn()
        measurements.append((time.perf_counter_ns()-start)/1000)
    return dict(median_us=float(np.median(measurements)),
                p90_us=float(np.percentile(measurements, 90)), raw_us=measurements)


def checks():
    maxerr = 0.
    p = {k:v.astype('float64') for k,v in initialize(68).items()}
    _,x,y = data(100, 1)
    _,g = loss_grad(p, x[:4], y[:4])
    for key in p:
        idx = tuple(0 for _ in p[key].shape)
        old = p[key][idx]
        p[key][idx] = old+1e-5
        plus = loss_grad(p,x[:4],y[:4])[0]
        p[key][idx] = old-1e-5
        minus = loss_grad(p,x[:4],y[:4])[0]
        p[key][idx] = old
        maxerr = max(maxerr, abs((plus-minus)/2e-5-g[key][idx]))
    assert maxerr < 1e-7
    for bits in (4,8):
        for w in (np.zeros((2,32),'float32'), np.linspace(-3,3,64,dtype='float32').reshape(2,32)):
            a=pack(w,bits,32)
            expected=a['q'].reshape(-1,32).astype('float32')*a['scales'][:,None]
            np.testing.assert_array_equal(unpack(a),expected.reshape(w.shape))
            assert len(a['payload']) == w.size*bits//8
    # Equal unweighted weight error can give radically different output error.
    x=np.diag([100.,1.]); e1=np.array([[.1,0.]]); e2=np.array([[0.,.1]])
    return dict(gradient_max_error=maxerr, packing_roundtrip='passed', zero_group='passed',
                equal_weight_mse=[float(np.mean(e*e)) for e in (e1,e2)],
                output_squared_error=[float(np.sum((e@x)**2)) for e in (e1,e2)])


def main():
    ap=argparse.ArgumentParser();ap.add_argument('--out',type=Path,default=Path(__file__).parent)
    args=ap.parse_args();args.out.mkdir(parents=True,exist_ok=True)
    datasets={k:data(seed,n) for k,seed,n in [('train',680,256),('dev',681,64),('test',682,128)]}
    seq_hashes={k:{hashlib.sha256(row.tobytes()).hexdigest() for row in v[0]} for k,v in datasets.items()}
    assert not (seq_hashes['train']&seq_hashes['dev'] or seq_hashes['train']&seq_hashes['test'] or seq_hashes['dev']&seq_hashes['test'])
    np.savez(args.out/'data.npz',**{k:v[0] for k,v in datasets.items()})
    results=dict(environment=dict(python=sys.version, numpy=np.__version__, platform=platform.platform(),
                                  thread_settings={k:os.environ[k] for k in ('OMP_NUM_THREADS','OPENBLAS_NUM_THREADS','VECLIB_MAXIMUM_THREADS')}),
                 checks=checks(), protocol=dict(train_sequences=256, dev_sequences=64,test_sequences=128,
                  tokens_per_sequence=32, context=2, vocab=16, calibration='none; weight absmax RTN',
                  bits=[8,4], group=32, bias_dtype='FP32', compute_dtype='FP32', timing='CPU forward batch1; no integer kernel; no GPU measurement'),runs=[])
    for seed in (68,69,70):
        p=initialize(seed); rng=np.random.default_rng(seed+1000)
        m={k:np.zeros_like(v) for k,v in p.items()};v={k:a.copy() for k,a in m.items()}
        _,tx,ty=datasets['train'];_,dx,dy=datasets['dev'];_,ex,ey=datasets['test']
        initial=float(nlls(p,dx,dy).mean());history=[];best=float('inf');chosen=None
        for step in range(1,801):
            idx=rng.integers(len(tx),size=128);loss,g=loss_grad(p,tx[idx],ty[idx])
            for k in p:
                m[k]=.9*m[k]+.1*g[k];v[k]=.999*v[k]+.001*g[k]*g[k]
                p[k]-=.01*(m[k]/(1-.9**step))/(np.sqrt(v[k]/(1-.999**step))+1e-8)
            if step in (200,400,600,800):
                dn=float(nlls(p,dx,dy).mean());history.append(dict(step=step,dev_nll=dn))
                if dn<best:best=dn;chosen={k:a.copy() for k,a in p.items()};beststep=step
        assert best<initial
        p=chosen;np.savez(args.out/f'fp32_seed{seed}.npz',**p)
        base=nlls(p,ex,ey);rows=[]
        configs=[('FP32',None,None),('INT8_G32',8,32),('INT4_TENSOR',4,None),('INT4_G32',4,32)]
        for name,bits,group in configs:
            folder=args.out/f'seed{seed}_{name}';folder.mkdir(exist_ok=True)
            q={} if bits is None else {k:pack(p[k],bits,group or p[k].size) for k in WEIGHTS}
            effective=restore(q,p);score=nlls(effective,ex,ey)
            np.save(folder/'token_nll.npy',score)
            manifest={}
            for k in p:
                if k in q:
                    a=q[k];(folder/f'{k}.bin').write_bytes(a['payload'])
                    (folder/f'{k}.scale').write_bytes(a['scales'].tobytes())
                    manifest[k]={z:a[z] for z in ('bits','group','shape')}
                    # Verify persisted integer payload/scales decode to the evaluated weights.
                    disk={**a,'payload':(folder/f'{k}.bin').read_bytes(),
                          'scales':np.frombuffer((folder/f'{k}.scale').read_bytes(),dtype='<f4')}
                    np.testing.assert_array_equal(unpack(disk),effective[k])
                else:(folder/f'{k}.f32').write_bytes(p[k].astype('<f4').tobytes())
            (folder/'manifest.json').write_text(json.dumps(manifest,indent=2))
            payload_bytes=sum(f.stat().st_size for f in folder.iterdir() if f.suffix in ('.bin','.scale','.f32'))
            cached=benchmark(lambda:forward(effective,ex[:1]))
            endtoend=cached if not q else benchmark(lambda:forward(restore(q,p),ex[:1]))
            row=dict(name=name,ppl=float(np.exp(score.mean())),nll=float(score.mean()),
                     delta_nll=float((score-base).mean()),payload_bytes=payload_bytes,
                     manifest_bytes=(folder/'manifest.json').stat().st_size,
                     fp32_materialized_bytes=sum(a.nbytes for a in effective.values()),
                     weight_mse=float(sum(np.sum((p[k]-effective[k])**2) for k in WEIGHTS)/sum(p[k].size for k in WEIGHTS)),
                     cached_forward=cached,decode_and_forward=endtoend)
            rows.append(row)
            print(seed,name,'PPL',round(row['ppl'],5),'payload',payload_bytes,
                  'cached_us',round(cached['median_us'],2),'decode_forward_us',round(endtoend['median_us'],2))
        results['runs'].append(dict(seed=seed,initial_dev_nll=initial,selected_step=beststep,history=history,results=rows))
    results['shapes']={k:list(v.shape) for k,v in p.items()}
    results['parameters']=sum(v.size for v in p.values())
    (args.out/'results.json').write_text(json.dumps(results,indent=2)+'\n')
    print('shapes',results['shapes'],'parameters',results['parameters'],'checks',results['checks'])

if __name__=='__main__': main()
