"""Recompute saved results with scalar routing independent of training forward."""
import json
import math
from pathlib import Path
import numpy as np

root = Path(__file__).parent / 'results'
data = json.loads((root/'data.json').read_text())
assert set(map(tuple,data['train']['x'])).isdisjoint(map(tuple,data['test']['x']))
assert set(map(tuple,data['train']['x'])).isdisjoint(map(tuple,data['dev']['x']))
assert set(map(tuple,data['dev']['x'])).isdisjoint(map(tuple,data['test']['x']))
summaries = {k: [] for k in ['task','balance','balance_z']}
for file in sorted(root.glob('*_7?.json')):
    run = json.loads(file.read_text())
    weights = json.loads(file.with_name(file.stem+'_weights.json').read_text())
    w, v = weights['w'], weights['v']
    for factor, expected in run['capacity_sweep'].items():
        means = {k:[] for k in ['mse','drop','max_share','load_cv','balance','z_loss','entropy']}
        for batch in range(4):
            cap = math.ceil(float(factor)*64/4)
            pre, post, losses, probs, zs, ent = [0]*4,[0]*4,[],[],[],[]
            for local in range(64):
                j=batch*64+local; x=data['test']['x'][j]
                logits=[sum(x[d]*w[d][e] for d in range(4)) for e in range(4)]
                maxz=max(logits); exps=[math.exp(z-maxz) for z in logits]
                p=[e/sum(exps) for e in exps]
                selected=max(range(4),key=lambda e:p[e])
                pre[selected]+=1; accepted=pre[selected]<=cap
                post[selected]+=int(accepted)
                pred=0.25*x[0]
                if accepted:
                    pred+=p[selected]*sum(x[d]*v[selected][d] for d in range(4))
                losses.append((pred-data['test']['y'][j])**2)
                probs.append(p); zs.append((maxz+math.log(sum(exps)))**2)
                ent.append(-sum(pe*math.log(pe) for pe in p))
                if factor=='1.25':
                    trace=run['traces'][batch]
                    assert trace['selected'][local]==selected and trace['accepted'][local]==accepted
                    assert np.allclose(trace['probabilities'][local],p,rtol=0,atol=1e-12)
                    assert abs(trace['predictions'][local]-pred)<1e-12
            f=np.array(pre)/64
            observed=dict(mse=np.mean(losses),drop=1-sum(post)/64,max_share=max(f),
                          load_cv=np.std(f)/np.mean(f),balance=4*np.dot(f,np.mean(probs,axis=0)),
                          z_loss=np.mean(zs),entropy=np.mean(ent))
            for k,val in observed.items():
                assert abs(expected['batches'][batch][k]-val)<1e-10,(file.name,factor,k)
                means[k].append(val)
            assert pre==expected['batches'][batch]['pre']
            assert post==expected['batches'][batch]['post']
        for k, values in means.items():
            assert abs(expected[k]-np.mean(values))<1e-10
    assert run['test']==run['capacity_sweep']['1.25']
    assert [h['step'] for h in run['history']]==list(range(0,601,100))
    summaries[run['name']].append(run['test'])
summary=json.loads((root/'summary.json').read_text())['results']
for name,runs in summaries.items():
    assert len(runs)==3
    for k,expected in summary[name].items():
        values=[r[k] for r in runs]
        assert abs(expected['mean']-np.mean(values))<1e-12
        assert abs(expected['sd']-np.std(values,ddof=1))<1e-12
print('PASS: 9 weights, 45 capacity evaluations, 11520 scalar token predictions, all first-capacity traces, 3-seed means/SD, disjoint splits.')
