"""Recompute saved embeddings and rankings without importing training code."""
import hashlib
import json
import math
from pathlib import Path
import numpy as np

ROOT = Path(__file__).parent/'results'
summary = json.loads((ROOT/'summary.json').read_text())
rows_checked, entries_checked, max_error = 0, 0, 0.
failures = []


def scalar_embeddings(x, w):
    result = []
    for row in x:
        z = [math.fsum(float(a)*float(b) for a,b in zip(row, col)) for col in w.T]
        n = math.sqrt(math.fsum(a*a for a in z))
        result.append([a/n for a in z])
    return result


def independent_metrics(s, qids, gids):
    ranks = [sorted(range(len(gids)), key=lambda j: (-row[j], j)) for row in s]
    out = {}
    for k in [1,3,5]:
        counts = [sum(gids[j] == q for j in rank[:k]) for q,rank in zip(qids,ranks)]
        denom = [sum(g == q for g in gids) for q in qids]
        out[f'hit@{k}'] = sum(c > 0 for c in counts)/len(qids)
        out[f'relevant_recall@{k}'] = sum(c/d for c,d in zip(counts,denom))/len(qids)
    out['mrr'] = sum(1/(next(i for i,j in enumerate(rank) if gids[j] == q)+1)
                     for q,rank in zip(qids,ranks))/len(qids)
    return out, ranks


for run in summary['runs']:
    seed = run['seed']
    data = np.load(ROOT/f'{seed}_data.npz')
    hashes = lambda x: {hashlib.sha256(row.tobytes()).hexdigest() for row in x}
    assert len(hashes(data['train_x'])) == 108
    assert len(hashes(data['test_x'])) == 45
    assert hashes(data['train_x']).isdisjoint(hashes(data['test_x']))
    assert np.array_equal(np.bincount(data['test_y']), np.full(9,5))
    for batch in data['batches']:
        assert np.array_equal(data['train_y'][batch], np.arange(9))
    for arm in ['baseline','aligned','mismatched']:
        state = data if arm == 'baseline' else np.load(ROOT/f'{seed}_{arm}.npz')
        wi, wt = (state['init_wi'],state['init_wt']) if arm == 'baseline' else (state['wi'],state['wt'])
        u,v = scalar_embeddings(data['test_x'],wi), scalar_embeddings(data['text'],wt)
        s = np.array([[math.fsum(a*b for a,b in zip(i,t)) for t in v] for i in u])
        if arm != 'baseline':
            max_error = max(max_error,float(np.max(np.abs(s-state['scores']))))
            np.testing.assert_allclose(s,state['scores'],atol=1e-12,rtol=0)
        entries_checked += s.size
        for direction,scores,qids,gids in [('i2t',s,list(data['test_y']),list(range(9))),
                                           ('t2i',s.T,list(range(9)),list(data['test_y']))]:
            metrics,ranks = independent_metrics(scores,qids,gids)
            for k,val in metrics.items():
                assert abs(val-run[arm][direction][k]) < 1e-12
            rows_checked += len(scores)
            if direction == 'i2t' and arm == 'aligned':
                for i,(qid,rank) in enumerate(zip(qids,ranks)):
                    pred = gids[rank[0]]
                    if pred != qid:
                        failures.append({'seed':seed,'test_row':i,'gold':int(qid),'prediction':int(pred),
                                         'same_color':bool(qid//3==pred//3),
                                         'same_shape':bool(qid%3==pred%3)})
for arm, directions in summary['aggregate'].items():
    for direction, stats in directions.items():
        vals=[r[arm][direction]['hit@1'] for r in summary['runs']]
        mean=sum(vals)/3
        sd=math.sqrt(sum((x-mean)**2 for x in vals)/2)
        assert abs(mean-stats['hit@1_mean']) < 1e-12
        assert abs(sd-stats['hit@1_sample_sd']) < 1e-12
report={'status':'passed','ranked_queries':rows_checked,'scalar_cosines':entries_checked,
        'max_score_error':max_error,'aligned_i2t_errors':len(failures),
        'same_color_errors':sum(x['same_color'] for x in failures),
        'same_shape_errors':sum(x['same_shape'] for x in failures),'failures':failures}
(ROOT/'audit.json').write_text(json.dumps(report,indent=2)+'\n')
print(json.dumps(report,indent=2))
