"""CPU-only protocol experiment. Frozen analytic scores, NOT trained neural nets."""
import json
import platform
from pathlib import Path
import numpy as np

ROOT = Path(__file__).resolve().parent
OUT = ROOT / 'results'
SEEDS = (78, 79, 80)
SPLITS = {'cal': 1024, 'policy': 1024, 'id': 4096, 'shift': 4096}
GRID = np.arange(50, 100, dtype=np.float64) / 100
TARGET_RISK = .10
MIN_ACCEPT = 100


def sigmoid(x):
    return np.exp(-np.logaddexp(0., -x))


def probabilities(z, t):
    assert np.isfinite(t) and t > 0
    a = z / t
    a = a - a.max(axis=1, keepdims=True)
    p = np.exp(a)
    return p / p.sum(axis=1, keepdims=True)


def bins(conf, correct, m):
    rows = []
    for j in range(m):
        mask = (conf > j / m) & (conf <= (j + 1) / m)
        n = int(mask.sum())
        rows.append({'lower': j/m, 'upper': (j+1)/m, 'n': n,
                     'confidence': float(conf[mask].mean()) if n else None,
                     'accuracy': float(correct[mask].mean()) if n else None})
    assert sum(r['n'] for r in rows) == len(conf)
    ece = sum(r['n'] / len(conf) * abs(r['accuracy'] - r['confidence'])
              for r in rows if r['n'])
    return float(ece), rows


def select(conf, correct, threshold):
    mask = conf >= threshold
    n = int(mask.sum())
    return {'threshold': float(threshold), 'accepted': n,
            'errors': int((~correct[mask]).sum()), 'coverage': n/len(conf),
            'risk': float((~correct[mask]).mean()) if n else None}


def metrics(z, y, t):
    p = probabilities(z, t)
    conf = p.max(1)
    correct = p.argmax(1) == y
    a = z/t
    nll = np.logaddexp.reduce(a, axis=1) - a[np.arange(len(y)), y]
    eces, reliability = {}, {}
    for m in (5, 10, 15, 30):
        eces[str(m)], reliability[str(m)] = bins(conf, correct, m)
    result = {'n': len(y), 'accuracy': float(correct.mean()),
              'mean_confidence': float(conf.mean()), 'nll': float(nll.mean()),
              'brier_sum_classes': float(((p-np.eye(2)[y])**2).sum(1).mean()),
              'ece': eces, 'reliability': reliability,
              'fixed_0.9': select(conf, correct, .9)}
    return result, p, conf, correct


def fit_temperature(z, y):
    # NLL is convex in inverse temperature beta. Solve derivative=0 on fixed interval.
    def grad(beta):
        p = probabilities(z, 1/beta)
        return float(((p*z).sum(1)-z[np.arange(len(y)), y]).mean())
    lo, hi = .05, 20.
    assert grad(lo) < 0 < grad(hi), 'Optimum not bracketed: report boundary instead.'
    trace = []
    for _ in range(80):
        mid = (lo+hi)/2
        g = grad(mid)
        trace.append([mid, g])
        if g > 0: hi = mid
        else: lo = mid
    beta = (lo+hi)/2
    return 1/beta, trace, grad(beta)


def policy(conf, correct):
    rows = [select(conf, correct, t) for t in GRID]
    feasible = [r for r in rows if r['accepted'] >= MIN_ACCEPT and r['risk'] <= TARGET_RISK]
    # Grid ascending; stable max picks lowest threshold among coverage ties.
    chosen = max(feasible, key=lambda r: r['coverage']) if feasible else select(conf, correct, 1.01)
    return chosen, rows


def self_checks():
    c = np.array([.6]*10 + [.9]*10)
    ok = np.array([True]*9+[False]+[True]*6+[False]*4)
    e1, _ = bins(c, ok, 1)
    e10, _ = bins(c, ok, 10)
    assert abs(e1) < 1e-14 and abs(e10-.3) < 1e-14
    assert bins(np.array([.5, 1.]), np.array([True, False]), 2)[1][0]['n'] == 1
    assert select(c, ok, 1.01)['risk'] is None
    assert select(np.array([.9]), np.array([False]), .9)['accepted'] == 1
    assert np.isfinite(probabilities(np.array([[10000., -10000.]]), 1)).all()
    multi = np.array([[2., 0., 0.], [1.5, 1.4, -10.]])
    c1 = probabilities(multi, 1).max(1)
    c10 = probabilities(multi, 10).max(1)
    assert c1[0] > c1[1] and c10[0] < c10[1]
    assert np.array_equal(probabilities(multi,1).argmax(1), probabilities(multi,10).argmax(1))
    return {'bin_cancellation': {'ece1':e1, 'ece10':e10},
            'multiclass_rank_reversal': {'logits':multi.tolist(),'T1':c1.tolist(),'T10':c10.tolist()},
            'empty_acceptance_risk':None, 'boundary_and_stability':'passed'}


def main():
    OUT.mkdir(exist_ok=True)
    summary = {'scope':'synthetic frozen analytic binary scorer; no neural training',
               'python':platform.python_version(), 'numpy':np.__version__,
               'seeds':list(SEEDS), 'split_sizes':SPLITS,
               'settings':{'threshold_grid':GRID.tolist(),'target_empirical_risk':TARGET_RISK,
                           'min_accepted_policy':MIN_ACCEPT,'temperature_bounds':[.05,20]},
               'self_checks':self_checks(), 'runs':[]}
    for seed in SEEDS:
        data = {}
        for k, (split,n) in enumerate(SPLITS.items()):
            rng = np.random.default_rng(np.random.SeedSequence([seed,k]))
            x = rng.normal(size=(n,3))
            s = x @ np.array([1.5,-.8,.5])
            truth = sigmoid(.5*s-1.) if split == 'shift' else sigmoid(s)
            y = (rng.random(n) < truth).astype(np.int64)
            z = np.column_stack([np.zeros(n),3*s])
            data[split] = (x,z,y,truth)
        allx = np.concatenate([v[0] for v in data.values()])
        assert len(np.unique(allx,axis=0)) == len(allx)
        t, trace, grad = fit_temperature(data['cal'][1],data['cal'][2])
        run = {'seed':seed,'temperature':t,'inverse_temperature_gradient':grad,
               'fit_trace':trace,'policies':{},'metrics':{},'binary_rank_invariant':{}}
        saved = {}
        for split,(x,z,y,truth) in data.items():
            for key,val in [('x',x),('z',z),('y',y),('truth',truth)]: saved[split+'_'+key]=val
        for arm,temp in [('raw',1.),('scaled',t)]:
            _,_,cc,oo = metrics(data['policy'][1],data['policy'][2],temp)
            chosen, sweep = policy(cc,oo)
            run['policies'][arm] = {'chosen':chosen,'sweep':sweep}
            run['metrics'][arm] = {}
            for split,(_,z,y,_) in data.items():
                result,p,conf,correct = metrics(z,y,temp)
                result['selected_policy'] = select(conf,correct,chosen['threshold'])
                run['metrics'][arm][split]=result
                saved[f'{split}_{arm}_probabilities']=p
                saved[f'{split}_{arm}_accepted']=conf>=chosen['threshold']
        for split in SPLITS:
            p0=saved[f'{split}_raw_probabilities']; p1=saved[f'{split}_scaled_probabilities']
            assert np.array_equal(p0.argmax(1),p1.argmax(1))
            order0=np.argsort(-p0.max(1),kind='stable'); order1=np.argsort(-p1.max(1),kind='stable')
            assert np.array_equal(order0,order1)
            correct=p0.argmax(1)==data[split][2]
            risk=np.cumsum(~correct[order0])/np.arange(1,len(correct)+1)
            saved[f'{split}_rank_order']=order0
            saved[f'{split}_risk_coverage']=np.column_stack([np.arange(1,len(correct)+1)/len(correct),risk])
            run['binary_rank_invariant'][split]=True
        assert run['metrics']['scaled']['cal']['nll'] <= run['metrics']['raw']['cal']['nll']
        np.savez_compressed(OUT/f'seed_{seed}.npz',**saved)
        summary['runs'].append(run)
        print(f'seed={seed} cal logits={data["cal"][1].shape} labels={data["cal"][2].shape} T={t:.9f} beta_gradient={grad:.3g}')
        for split in ('id','shift'):
            for arm in ('raw','scaled'):
                r=run['metrics'][arm][split]; q=r['selected_policy']; f=r['fixed_0.9']
                print(f'  {split:5} {arm:6} acc={r["accuracy"]:.6f} NLL={r["nll"]:.6f} ECE15={r["ece"]["15"]:.6f} fixed90 cov={f["coverage"]:.6f} risk={f["risk"]} policy={q}')
    aggregate={}
    for split in ('id','shift'):
        aggregate[split]={}
        for arm in ('raw','scaled'):
            rows=[r['metrics'][arm][split] for r in summary['runs']]
            cols={k:[r[k] for r in rows] for k in ('accuracy','nll','brier_sum_classes')}
            cols['ece15']=[r['ece']['15'] for r in rows]
            for kind in ('fixed_0.9','selected_policy'):
                for k in ('risk','coverage'):
                    cols[kind+'_'+k]=[r[kind][k] for r in rows]
            aggregate[split][arm]={k:{'mean':float(np.mean(v)) if all(x is not None for x in v) else None,
                                    'seed_sd':float(np.std(v,ddof=1)) if all(x is not None for x in v) else None,
                                    'values':v,'defined_seeds':sum(x is not None for x in v)} for k,v in cols.items()}
    summary['aggregate']=aggregate
    (OUT/'summary.json').write_text(json.dumps(summary,ensure_ascii=False,indent=2,allow_nan=False)+'\n')
    print('PASS: split uniqueness, positive T, calibration NLL, argmax, binary ranking, bin edges, cancellation, empty risk, numerical stability.')


if __name__ == '__main__':
    main()
