"""Independent scalar audit: does not import experiment.py."""
from pathlib import Path
import json, math
import numpy as np
P=Path(__file__).parent/'results'

def stats(golds,preds):
    tp=sum(len(g&p) for g,p in zip(golds,preds));fp=sum(len(p-g) for g,p in zip(golds,preds));fn=sum(len(g-p) for g,p in zip(golds,preds))
    return dict(tp=tp,fp=fp,fn=fn,micro_f1=2*tp/max(1,2*tp+fp+fn),set_em=sum(g==p for g,p in zip(golds,preds))/len(golds),empty_queries=sum(not g for g in golds))

def gates(w):
    return [[math.exp(w[2*r+c])/sum(math.exp(w[2*r+j]) for j in range(2)) for c in range(2)] for r in range(2)]

def solve(edges,s,q,candidates):
    return {v for u,r,m in edges for mm,rr,v in edges if u==s and r==q[0] and mm==m and rr==q[1] and v in candidates}

def main():
    summary=json.loads((P/'summary.json').read_text());seen=set();max_error=0.;queries=0;logs=0;witness=None
    for run in summary['runs']:
        seed=run['seed']
        for split in ['train','dev','test']:
            d=np.load(P/f'{seed}_{split}.npz');a,s,q,y=[d[k] for k in ['a','s','q','y']]
            for i in range(len(a)):
                key=a[i].astype('uint8').tobytes();assert key not in seen;seen.add(key)
                edges={(int(u),int(r),int(v)) for r,u,v in np.argwhere(a[i])}
                gold=solve(edges,int(s[i]),q[i],set(range(12)))
                assert gold==set(np.flatnonzero(y[i]));queries+=1
        model=np.load(P/f'{seed}_model.npz');p=model['weights'];w1,w2=gates(p[:4]),gates(p[4:8])
        golds=[set(map(int,np.flatnonzero(row))) for row in y]
        for mode in ['intact','reversed','no_edges']:
            aa=a if mode=='intact' else a.transpose(0,1,3,2) if mode=='reversed' else np.zeros_like(a)
            preds=[];loss=0.
            for i in range(len(a)):
                h=[sum(w1[q[i,0]][r]*aa[i,r,s[i],u] for r in range(2)) for u in range(12)]
                z=[math.exp(p[8])*sum(w2[q[i,1]][r]*h[u]*aa[i,r,u,v] for r in range(2) for u in range(12))+p[9] for v in range(12)]
                max_error=max(max_error,max(abs(float(model[mode+'_logits'][i,v])-z[v]) for v in range(12)))
                preds.append({v for v in range(12) if z[v]>=0})
                loss+=sum(max(t,0)+math.log1p(math.exp(-abs(t)))-y[i,v]*t for v,t in enumerate(z))/y.size
            measured=stats(golds,preds);measured['bce']=loss
            for key,val in measured.items():assert abs(val-run['controls'][mode][key])<1e-10
        for k in [1,3,6,12]:
            for mode in ['induced','bridge']:
                rows=json.loads((P/f'{seed}_{mode}_k{k}.json').read_text());preds=[];recalls=[];n_edges=[]
                for i,row in enumerate(rows):
                    prompt=row['prompt'];assert set(prompt)=={'source','relations','allowed_answers','triples'}
                    assert row['query_id']==i and prompt['source']==s[i] and prompt['relations']==q[i].tolist()
                    candidates=sorted(range(12),key=lambda v:(-float(model['logits'][i,v]),v))[:k]
                    assert prompt['allowed_answers']==candidates
                    all_edges={(int(u),int(r),int(v)) for r,u,v in np.argwhere(a[i])}
                    allowed=set(candidates)|{int(s[i])}
                    expected={t for t in all_edges if ((t[0] in allowed and t[2] in allowed) if mode=='induced' else (t[0]==s[i] or t[2] in candidates))}
                    actual={tuple(t) for t in prompt['triples']};assert actual==expected
                    pred=solve(actual,int(s[i]),q[i],set(candidates));assert pred==set(row['answer'])
                    assert set(row['gold'])==golds[i] and pred<=golds[i]
                    preds.append(pred);n_edges.append(len(actual));logs+=1
                    if golds[i]:recalls.append(len(set(candidates)&golds[i])/len(golds[i]))
                    if witness is None and mode=='induced' and k==3 and (set(candidates)&golds[i])-pred:
                        target=min((set(candidates)&golds[i])-pred)
                        paths=[[[int(s[i]),int(q[i,0]),m],[m,int(q[i,1]),target]] for m in range(12) if a[i,q[i,0],s[i],m] and a[i,q[i,1],m,target]]
                        witness={'seed':seed,'query_id':i,'source':int(s[i]),'relations':q[i].tolist(),'candidates':candidates,'gold':sorted(golds[i]),'induced_answer':sorted(pred),'missing_target':target,'full_graph_paths':paths}
                measured=stats(golds,preds);measured.update(nonempty_candidate_recall=sum(recalls)/len(recalls),mean_triples=sum(n_edges)/len(n_edges))
                for key,val in measured.items():assert abs(val-run['budgets'][f'{mode}_k{k}'][key])<1e-10
        # Recompute reported mean and sample SD independently.
    for mode,metrics in summary['aggregate'].items():
        for metric,values in metrics.items():
            xs=[r['controls'][mode][metric] for r in summary['runs']];mean=sum(xs)/3;sd=math.sqrt(sum((x-mean)**2 for x in xs)/2)
            assert abs(mean-values['mean'])<1e-12 and abs(sd-values['seed_sd'])<1e-12
    assert max_error<1e-10
    audit={'status':'passed','unique_graphs':len(seen),'ground_truth_queries_rebuilt':queries,'logit_values_rebuilt':3*3*128*12,'evidence_records_replayed':logs,'max_logit_abs_error':max_error,'witness':witness}
    (P/'audit.json').write_text(json.dumps(audit,indent=2));print(json.dumps(audit,indent=2))
if __name__=='__main__':main()
