"""Independent scalar replay; does not import the training script."""
import json
import math
import statistics
from pathlib import Path
import numpy as np

root=Path(__file__).parent/'results'
s=json.loads((root/'summary.json').read_text())
transitions=predictions=metrics=0
max_error=0.
initials=set()
for seed in s['seeds']:
    d=np.load(root/f'data_{seed}.npz');w=np.load(root/f'weights_{seed}.npz');p=np.load(root/f'predictions_{seed}.npz')
    for split,n in s['split_counts'].items():
        states=d[f'{split}_states'];actions=d[f'{split}_actions']
        assert states.shape==(n,41,2) and actions.shape==(n,40)
        for st,ac in zip(states,actions):
            pair=tuple(st[0]);assert pair not in initials;initials.add(pair)
            for t,a in enumerate(ac):
                x,v=map(float,st[t]);vn=.98*v-.1*math.sin(x)+.08*float(a)
                assert math.isclose(vn,float(st[t+1,1]),abs_tol=1e-13)
                assert math.isclose(x+.1*vn,float(st[t+1,0]),abs_tol=1e-13)
                transitions+=1
    for row in [r for r in s['runs'] if r['seed']==seed]:
        mode=row['mode'];weights=w[mode].tolist()
        if mode=='sine_state':
            expected=[[1,0],[.098,.98],[.008,.08],[0,0],[-.01,-.1]]
            assert np.max(np.abs(w[mode]-expected))<1e-10
        for split in ('dev','id','ood'):
            truth=d[f'{split}_states'].tolist();actions=d[f'{split}_actions'].tolist()
            for condition in ('teacher','open','correct5'):
                saved=p[f'{mode}_{split}_{condition}'];sqx=[[] for _ in range(40)];sqv=[[] for _ in range(40)]
                for i,(st,ac) in enumerate(zip(truth,actions)):
                    current=st[0]
                    for t,a in enumerate(ac):
                        if condition=='teacher' or (condition=='correct5' and t%5==0): current=st[t]
                        x,v=current
                        f=[x,a,1.] if mode=='position_only' else [x,v,a,1.]
                        if mode=='sine_state':f.append(math.sin(x))
                        nxt=[sum(f[k]*weights[k][j] for k in range(len(f))) for j in range(2)]
                        idx=t if condition=='teacher' else t+1
                        delta=max(abs(nxt[j]-float(saved[i,idx,j])) for j in range(2))
                        max_error=max(max_error,delta);assert delta<1e-10
                        sqx[t].append((nxt[0]-st[t+1][0])**2);sqv[t].append((nxt[1]-st[t+1][1])**2)
                        current=nxt;predictions+=1
                xs=[statistics.mean(v) for v in sqx];vs=[statistics.mean(v) for v in sqv];r=row['metrics'][split]
                pairs=[(statistics.mean(xs),r['teacher_position_mse']),(statistics.mean(vs),r['teacher_velocity_mse'])] if condition=='teacher' else list(zip(xs,r[condition]['position_mse_by_horizon']))+list(zip(vs,r[condition]['velocity_mse_by_horizon']))+[(statistics.mean(xs),r[condition]['position_mse_all'])]
                for a,b in pairs:
                    assert math.isclose(a,b,abs_tol=1e-10,rel_tol=1e-9);metrics+=1
for mode,ds in s['aggregate'].items():
    for split,agg in ds.items():
        rr=[r['metrics'][split] for r in s['runs'] if r['mode']==mode]
        for name,value in agg.items():
            if name=='teacher_x':vals=[r['teacher_position_mse'] for r in rr]
            elif name=='open_all':vals=[r['open']['position_mse_all'] for r in rr]
            elif name=='correct5_all':vals=[r['correct5']['position_mse_all'] for r in rr]
            else:vals=[r['open']['position_mse_by_horizon'][int(name.split('h')[1])-1] for r in rr]
            for a,b in [(statistics.mean(vals),value['mean']),(statistics.stdev(vals),value['seed_sd'])]:
                assert math.isclose(a,b,abs_tol=1e-12);metrics+=1
result={'status':'passed','true_transitions_replayed':transitions,'predicted_state_vectors_replayed':predictions,'metric_values_recomputed':metrics,'unique_initial_states':len(initials),'max_prediction_abs_difference':max_error,'sine_coefficients_match_true_dynamics':True}
(root/'audit.json').write_text(json.dumps(result,indent=2)+'\n');print(json.dumps(result,indent=2))
