"""Independent scalar forward pass and saved-result audit. Does not import trainer."""
import hashlib
import json
import math
from pathlib import Path
import numpy as np

ROOT=Path(__file__).resolve().parent
D=dict(np.load(ROOT/'results/data.npz'))
S=json.loads((ROOT/'results/summary.json').read_text())
max_error=0.;checked=0;cache={}

def compare(a,b):
    global max_error
    e=float(np.max(np.abs(np.asarray(a)-np.asarray(b))))
    max_error=max(max_error,e)
    assert e<1e-9,e

def dense(row,w,b):
    return [sum(float(row[i])*float(w[i,j]) for i in range(len(row)))+float(b[j]) for j in range(len(b))]

def fw(x,p):
    hs=[[],[],[]];zs=[]
    for row in x:
        a=[math.tanh(t) for t in dense(row,p['W1'],p['b1'])]
        b=[math.tanh(t) for t in dense(a,p['W2'],p['b2'])]
        hs[0].append(row);hs[1].append(a);hs[2].append(b)
        zs.append(dense(b,p['W3'],p['b3']))
    return list(map(np.asarray,hs)),np.asarray(zs)

def score(h,f,w):
    return np.asarray([sum((float(v)-float(m))/float(s)*float(a) for v,m,s,a in zip(row,f['mean'],f['scale'],w[:-1]))+float(w[-1]) for row in h])

def nll(z,y):
    return sum(max(float(t),0)+math.log1p(math.exp(-abs(float(t))))-int(q)*float(t) for t,q in zip(z,y))/len(y)

for split in ['train','dev','test']:
    x=D[split+'_x'];pos=(x[:,2]>0).astype(int)
    assert np.array_equal(D[split+'_task'],(x[np.arange(len(x)),pos]>0).astype(int))
    assert np.array_equal(D[split+'_unused'],(x[np.arange(len(x)),1-pos]>0).astype(int))
    assert D[split+'_shuffle'].sum()==len(x)/2
for seed in S['model_seeds']:
    path=ROOT/'checkpoints'/f'model_{seed}.npz'
    assert hashlib.sha256(path.read_bytes()).hexdigest()==S['checkpoint_hashes'][path.name]
    p=dict(np.load(path))
    cache[seed]={s:fw(D[s+'_x'],p) for s in ['train','dev','test']}
    saved=np.load(ROOT/'results'/f'behavior_{seed}.npz')
    compare(cache[seed]['test'][1],saved['original'])
    x=D['test_x'].copy();pos=(x[:,2]>0).astype(int);x[np.arange(len(x)),1-pos]*=-1
    compare(fw(x,p)[1],saved['flipped'])
    basis=saved['basis'];compare(basis.T@basis,np.eye(basis.shape[1]))
    compare(basis@basis.T@p['W3'],p['W3'])
    expected=cache[seed]['test'][0][2]@basis@basis.T
    compare(expected,saved['projected_h2'])
    compare([dense(row,p['W3'],p['b3']) for row in expected],saved['projected'])
    probe=np.load(ROOT/'results'/f's{seed}_l2_unused.npz')
    compare(score(expected,probe,probe['w']),saved['projected_probe'])
    b=next(r for r in S['behavior'] if r['seed']==seed)
    compare(float(np.mean(saved['original'].argmax(1)==D['test_task'])),b['original_accuracy'])
    compare(float(np.mean(saved['original'].argmax(1)!=saved['flipped'].argmax(1))),b['unused_flip_answer_change'])
    compare(float(np.max(np.abs(saved['original']-saved['flipped']))),b['unused_flip_max_logit_change'])
    compare(float(np.mean((saved['projected_probe']>=0)==D['test_unused'])),b['projected_unused_probe_accuracy'])
for rec in S['records']:
    f=np.load(ROOT/'results'/rec['file']);hs=cache[rec['seed']];layer=rec['layer'];target=rec['target']
    compare(f['mean'],hs['train'][0][layer].mean(0))
    compare(f['scale'],np.maximum(hs['train'][0][layer].std(0),1e-8))
    losses=[nll(score(hs['dev'][0][layer],f,w),D['dev_'+target]) for w in f['candidate_w']]
    compare(losses,rec['dev_candidate_nll']);assert int(np.argmin(losses))==rec['chosen_index']
    assert S['lambdas'][rec['chosen_index']]==rec['lambda_']
    compare(f['w'],f['candidate_w'][rec['chosen_index']])
    for split in ['train','dev','test']:
        z=score(hs[split][0][layer],f,f['w']);compare(z,f[split+'_logit'])
        compare(float(np.mean((z>=0)==D[split+'_'+target])),rec[split]['accuracy'])
        compare(nll(z,D[split+'_'+target]),rec[split]['nll']);checked+=len(z)
for key,v in S['aggregate'].items():
    layer=int(key[1]);target=key[3:]
    vals=[r['test']['accuracy'] for r in S['records'] if r['layer']==layer and r['target']==target]
    compare(np.mean(vals),v['mean']);compare(np.std(vals,ddof=1),v['sd'])
report=dict(status='passed',scalar_probe_predictions=checked,max_absolute_error=max_error,
            checks=['scalar backbone forward','scalar probe logits/NLL','train-only normalization','dev-only lambda choice',
                    'all aggregate means/sample SD','counterfactual inputs','projection invariance','checkpoint SHA256'])
(ROOT/'results/audit.json').write_text(json.dumps(report,indent=2)+'\n')
print(json.dumps(report,indent=2))
