"""Independent scalar forward replay; does not import experiment functions."""
import json
import math
from pathlib import Path
import numpy as np

ROOT=Path(__file__).parent/'results'


def affine(x,w,b):
    return [sum(x[k]*float(w[k,j]) for k in range(len(x)))+float(b[j]) for j in range(len(b))]


def forward(x,p):
    a=[math.tanh(v) for v in affine(x,p['W1'],p['b1'])]
    b=[math.tanh(v) for v in affine(a,p['W2'],p['b2'])]
    return affine(b,p['W3'],p['b3']),[a,b]


def metric(z,y): return z[y]-z[1-y]


def main():
    data=np.load(ROOT/'data.npz'); x=data['eval_x'];y=data['eval_y']
    summary=json.loads((ROOT/'summary.json').read_text()); maxerr=0.;count=0; diagnostics=[]
    for run in summary['runs']:
        seed=run['seed'];p=np.load(ROOT/f'model_{seed}.npz');cache={};baseline={}
        clean=[forward(v,p) for v in x]
        for condition in ['selected_flip','unselected_flip']:
            bx=x.copy()
            for i,row in enumerate(bx):
                j=1 if row[2]>0 else 0
                if condition=='unselected_flip':j=1-j
                row[j]=-row[j]
            bad=[forward(v,p) for v in bx]
            baseline[condition]=bad
            for layer in [1,2]:
                for basis in (['native','rotated'] if layer==2 else ['native']):
                    q=p['rotation'] if basis=='rotated' else np.eye(16)
                    cc=[affine(v[1][layer-1],q,[0]*16) for v in clean]
                    bb=[affine(v[1][layer-1],q,[0]*16) for v in bad]
                    cache[condition,layer,basis]=(cc,bb,q)
        rows=json.loads((ROOT/f'records_{seed}.json').read_text()); groups={}
        for r in rows:
            i=r['sample'];layer=r['layer'];cc,bb,q=cache[r['condition'],layer,r['basis']]
            h=(cc[i] if r['source']=='reverse' else bb[i]).copy()
            for j in r['indices']:
                h[j]={'clean':cc[i][j],'self':bb[i][j],'reverse':bb[i][j],
                      'zero':0.,'shuffle':cc[(i-1)%len(x)][j]}[r['source']]
            h=affine(h,q.T,[0]*16)
            if layer==1:h=[math.tanh(v) for v in affine(h,p['W2'],p['b2'])]
            z=affine(h,p['W3'],p['b3']);ld=metric(z,int(y[i]));prob=1/(1+math.exp(-ld))
            cl=metric(clean[i][0],int(y[i]));bl=metric(baseline[r['condition']][i][0],int(y[i]))
            values=[(z[j],r['logits'][j]) for j in [0,1]]+[(ld,r['ld']),(prob,r['prob']),
                    (cl,r['clean_ld']),(bl,r['corrupt_ld']),(ld-bl,r['delta_ld']),
                    (prob-1/(1+math.exp(-bl)),r['delta_prob'])]
            if r['recovery'] is not None:
                assert r['condition']=='selected_flip' and cl-bl>1e-6
                values.append(((ld-bl)/(cl-bl),r['recovery']))
            else:assert r['condition']=='unselected_flip'
            if r['source']=='reverse':values.append((cl-ld,r['degradation']))
            maxerr=max(maxerr,max(abs(a-b) for a,b in values));count+=1
            key=(r['condition'],layer,r['basis'],r['intervention'])
            groups.setdefault(key,[]).append(r)
        for g in run['groups']:
            key=tuple(g[k] for k in ['condition','layer','basis','intervention']);rr=groups[key]
            assert len(rr)==g['n']
            for raw,field in [('delta_ld','mean_delta_ld'),('delta_prob','mean_delta_prob')]:
                assert abs(sum(r[raw] for r in rr)/len(rr)-g[field])<1e-12
            assert sum(r['ld']>0 for r in rr)/len(rr)==g['accuracy']
            if g['mean_recovery'] is not None:
                assert abs(sum(r['recovery'] for r in rr)/len(rr)-g['mean_recovery'])<1e-12
        # Independently choose neuron sets using dev paired data only.
        dx=data['dev_x'];dy=data['dev_y'];dc=[forward(v,p) for v in dx];db=[]
        for row in dx:
            b=row.copy();j=int(b[2]>0);b[j]*=-1;db.append(forward(b,p))
        for layer in [1,2]:
            effects=[]
            for j in range(16):
                total=0.
                for i in range(len(dx)):
                    h=db[i][1][layer-1].copy();h[j]=dc[i][1][layer-1][j]
                    if layer==1:h=[math.tanh(v) for v in affine(h,p['W2'],p['b2'])]
                    z=affine(h,p['W3'],p['b3'])
                    total+=metric(z,int(dy[i]))-metric(db[i][0],int(dy[i]))
                effects.append(total/len(dx))
            assert sorted(run['top4'][str(layer)])==sorted(sorted(range(16),key=lambda j:effects[j])[-4:])
        def rows_for(layer,basis,name):return groups['selected_flip',layer,basis,name]
        residuals={}
        for layer in [1,2]:
            top=run['top4'][str(layer)];joint=rows_for(layer,'native','top4')
            residuals[str(layer)]=max(abs(joint[i]['delta_ld']-sum(rows_for(layer,'native','single_'+str(j))[i]['delta_ld'] for j in top)) for i in range(64))
        assert residuals['2']<1e-10
        basis_effects={}
        for basis in ['native','rotated']:
            effects=[sum(r['delta_ld'] for r in rows_for(2,basis,'single_'+str(j)))/64 for j in range(16)]
            basis_effects[basis]=dict(peak=max(effects),peak_index=int(np.argmax(effects)),sum=sum(effects),effects=effects)
        assert abs(basis_effects['native']['sum']-basis_effects['rotated']['sum'])<1e-10
        diagnostics.append(dict(seed=seed,nonadditivity_max_abs=residuals,basis=basis_effects))
    assert maxerr<1e-10,maxerr
    report=dict(status='passed',records_replayed=count,max_abs_error=maxerr,diagnostics=diagnostics,
                checks=['scalar forward and intervention replay','raw and normalized metrics',
                        'dev-only top4 selection','summary aggregates','linear readout additivity'])
    (ROOT/'audit.json').write_text(json.dumps(report,indent=2)+'\n')
    print(json.dumps(report,indent=2))


if __name__=='__main__':main()
