"""CPU relation-gated message passing; deterministic evidence reader, NOT an LLM."""
from pathlib import Path
import json, platform
import numpy as np
OUT = Path(__file__).parent / 'results'

def softmax(x):
    e = np.exp(x - x.max(axis=-1, keepdims=True)); return e / e.sum(axis=-1, keepdims=True)

def generate(seed, count, n=12):
    rng = np.random.default_rng(seed)
    a = (rng.random((count, 2, n, n)) < .12).astype(float)
    a[:, :, np.arange(n), np.arange(n)] = 0
    s = rng.integers(n, size=count); q = rng.integers(2, size=(count, 2))
    y = np.zeros((count, n))
    for i in range(count):
        # Ground truth uses explicit set traversal, independent of GNN matrix product.
        for m in range(n):
            if a[i, q[i, 0], s[i], m]:
                for v in range(n):
                    if a[i, q[i, 1], m, v]: y[i, v] = 1
    return a, s, q, y

def forward(p, data, gradient=False):
    a, s, q, y = data; b = len(a); ids = np.arange(b)
    w1 = softmax(p[:4].reshape(2, 2)); w2 = softmax(p[4:8].reshape(2, 2))
    c1, c2 = w1[q[:, 0]], w2[q[:, 1]]
    first = a[ids, :, s, :]
    h1 = np.einsum('br,brn->bn', c1, first)
    per_relation = np.einsum('bu,bruv->brv', h1, a)
    h2 = np.einsum('br,brv->bv', c2, per_relation)
    scale = np.exp(p[8]); z = scale*h2 + p[9]
    loss = np.mean(np.logaddexp(0, z)-y*z)
    if not gradient: return z, float(loss)
    dz = (1/(1+np.exp(-z))-y)/y.size
    dh2 = dz*scale
    dc2 = np.einsum('bv,brv->br', dh2, per_relation)
    dh1 = np.einsum('bv,br,bruv->bu', dh2, c2, a)
    dc1 = np.einsum('bn,brn->br', dh1, first)
    g = np.zeros(10)
    for offset, c, dc, rel in [(0,c1,dc1,q[:,0]), (4,c2,dc2,q[:,1])]:
        dw = c*(dc-(dc*c).sum(axis=1,keepdims=True)); dest = g[offset:offset+4].reshape(2,2)
        np.add.at(dest,rel,dw)
    g[8] = np.sum(dz*scale*h2); g[9] = dz.sum()
    return z, float(loss), g

def metrics(y, pred):
    yt, yp = y.astype(bool), pred.astype(bool)
    tp = int((yt & yp).sum()); fp=int((~yt & yp).sum()); fn=int((yt & ~yp).sum())
    return {'micro_f1':2*tp/max(1,2*tp+fp+fn), 'set_em':float((yt==yp).all(axis=1).mean()),
            'tp':tp,'fp':fp,'fn':fn, 'empty_queries':int((~yt.any(axis=1)).sum())}

def reader(triples, source, relations, candidates):
    mids={v for u,r,v in triples if u==source and r==relations[0]}
    return sorted({v for u,r,v in triples if u in mids and r==relations[1] and v in candidates})

def evidence_rows(a,s,q,y,z,k,mode):
    rows=[]; pred=np.zeros_like(y); recalls=[]
    for i in range(len(a)):
        candidates=np.argsort(-z[i],kind='stable')[:k].tolist()
        allowed={int(s[i]),*candidates}
        triples=[[int(u),int(r),int(v)] for r,u,v in np.argwhere(a[i])
                 if ((u in allowed and v in allowed) if mode=='induced' else (u==s[i] or v in candidates))]
        answer=reader(triples,int(s[i]),q[i].tolist(),candidates)
        pred[i,answer]=1
        gold=np.flatnonzero(y[i]).tolist()
        if gold: recalls.append(len(set(candidates)&set(gold))/len(gold))
        # Prompt does not contain gold, probabilities or labels. IDs are local to this graph.
        prompt={'source':int(s[i]),'relations':q[i].tolist(),'allowed_answers':candidates,'triples':triples}
        rows.append({'query_id':i,'prompt':prompt,'answer':answer,'gold':gold})
    result=metrics(y,pred)
    result.update({'nonempty_candidate_recall':float(np.mean(recalls)),
                   'mean_triples':float(np.mean([len(x['prompt']['triples']) for x in rows]))})
    return rows,result

def main():
    OUT.mkdir(exist_ok=True); runs=[]; all_hashes=set()
    for seed in [75,76,77]:
        train=generate(seed*10,256); dev=generate(seed*10+1,64); test=generate(seed*10+2,128)
        for split, data in [('train',train),('dev',dev),('test',test)]:
            a,s,q,y=data
            for graph in a:
                key=graph.astype('uint8').tobytes(); assert key not in all_hashes; all_hashes.add(key)
            np.savez_compressed(OUT/f'{seed}_{split}.npz',a=a,s=s,q=q,y=y)
        rng=np.random.default_rng(seed); p=np.r_[rng.normal(0,.05,8),0.,-2.]
        _,initial,g=forward(p,train,True); eps=1e-5; errs=[]
        for j in range(10):
            plus=p.copy();minus=p.copy();plus[j]+=eps;minus[j]-=eps
            numeric=(forward(plus,train)[1]-forward(minus,train)[1])/(2*eps)
            errs.append(abs(numeric-g[j]))
        assert max(errs)<1e-7
        m=np.zeros(10);v=np.zeros(10);curve=[]
        for step in range(1,501):
            _,loss,g=forward(p,train,True)
            m=.9*m+.1*g;v=.999*v+.001*g*g
            p-=.03*(m/(1-.9**step))/(np.sqrt(v/(1-.999**step))+1e-8)
            if step==1 or step%50==0: curve.append({'step':step,'train_loss':forward(p,train)[1],'dev_loss':forward(p,dev)[1]})
        a,s,q,y=test; z,_=forward(p,test)
        controls={}
        arrays={'weights':p,'logits':z}
        for name, aa in [('intact',a),('reversed',a.transpose(0,1,3,2)),('no_edges',np.zeros_like(a))]:
            zz,loss=forward(p,(aa,s,q,y));assert np.isfinite(zz).all()
            controls[name]={**metrics(y,zz>=0),'bce':loss};arrays[name+'_logits']=zz
        # Consistent relabeling must preserve logits (tie-breaking of top-k is not claimed invariant).
        perm=rng.permutation(12); inverse=np.argsort(perm)
        zp,_=forward(p,(a[:,:,perm,:][:,:,:,perm],inverse[s],q,y[:,perm]))
        permutation_error=float(np.max(abs(zp-z[:,perm])));assert permutation_error<1e-10
        budgets={}
        for k in [1,3,6,12]:
            for mode in ['induced','bridge']:
                rows,stats=evidence_rows(a,s,q,y,z,k,mode)
                budgets[f'{mode}_k{k}']=stats
                (OUT/f'{seed}_{mode}_k{k}.json').write_text(json.dumps(rows,indent=2))
        assert budgets['bridge_k12']['set_em']==1
        assert budgets['bridge_k3']['fp']==0
        np.savez_compressed(OUT/f'{seed}_model.npz',**arrays)
        run={'seed':seed,'initial_train_bce':initial,'final_train_bce':forward(p,train)[1],
             'gradient_max_abs_error':max(errs),'permutation_max_abs_error':permutation_error,
             'relation_gates':[softmax(p[:4].reshape(2,2)).tolist(),softmax(p[4:8].reshape(2,2)).tolist()],
             'controls':controls,'budgets':budgets,'curve':curve}
        runs.append(run)
    aggregate={}
    for name in runs[0]['controls']:
        aggregate[name]={metric:{'mean':float(np.mean(vals:=[r['controls'][name][metric] for r in runs])),
                               'seed_sd':float(np.std(vals,ddof=1))} for metric in ['micro_f1','set_em','bce']}
    summary={'python':platform.python_version(),'numpy':np.__version__,'parameters':10,'steps':500,'lr':.03,
             'seeds':[75,76,77],'split_sizes':[256,64,128],'a_shape':[256,2,12,12],
             'logits_shape':[256,12],'unique_graphs':len(all_hashes),'runs':runs,'aggregate':aggregate,
             'scope':'Trained toy message-passing GNN and deterministic reader; no LLM, PCST or paper benchmark executed.'}
    (OUT/'summary.json').write_text(json.dumps(summary,indent=2))
    print(json.dumps({k:summary[k] for k in ['python','numpy','parameters','a_shape','logits_shape','unique_graphs','aggregate']},indent=2))
    for r in runs: print('seed',r['seed'],'gradient',r['gradient_max_abs_error'],'budget',json.dumps(r['budgets']))
if __name__=='__main__':main()
