"""Trained toy MLP activation patching. No downloaded data or model weights."""
import argparse
import itertools
import json
import platform
from pathlib import Path
import numpy as np


def dataset(rng, repeats):
    core = np.repeat(np.array(list(itertools.product([-1., 1.], repeat=3))), repeats, axis=0)
    x = np.c_[core, rng.normal(0, .3, (len(core), 2))]
    selected = (x[:, 2] > 0).astype(int)
    y = (x[np.arange(len(x)), selected] > 0).astype(int)
    return x, y


def corrupt(x, relevant=True):
    z = x.copy()
    pos = (x[:, 2] > 0).astype(int)
    if not relevant:
        pos = 1 - pos
    z[np.arange(len(z)), pos] *= -1
    return z


def init(seed):
    r = np.random.default_rng(seed)
    p = {}
    for i, (a, b) in enumerate([(5, 16), (16, 16), (16, 2)], 1):
        p['W'+str(i)] = r.normal(0, 1/np.sqrt(a), (a, b))
        p['b'+str(i)] = np.zeros(b)
    return p


def forward(x, p):
    h1 = np.tanh(x @ p['W1'] + p['b1'])
    h2 = np.tanh(h1 @ p['W2'] + p['b2'])
    return h2 @ p['W3'] + p['b3'], [h1, h2]


def resume(h, layer, p):
    if layer == 1:
        h = np.tanh(h @ p['W2'] + p['b2'])
    return h @ p['W3'] + p['b3']


def loss_grad(x, y, p):
    z, (h1, h2) = forward(x, p)
    shifted = z-z.max(1, keepdims=True)
    e = np.exp(shifted); prob = e/e.sum(1, keepdims=True)
    loss = np.mean(np.log(e.sum(1))-shifted[np.arange(len(y)), y])
    dz = prob.copy(); dz[np.arange(len(y)), y] -= 1; dz /= len(y)
    d2 = (dz @ p['W3'].T)*(1-h2*h2)
    d1 = (d2 @ p['W2'].T)*(1-h1*h1)
    return float(loss), dict(W1=x.T@d1, b1=d1.sum(0), W2=h1.T@d2,
                             b2=d2.sum(0), W3=h2.T@dz, b3=dz.sum(0))


def scores(z, y):
    ld = z[np.arange(len(y)), y]-z[np.arange(len(y)), 1-y]
    prob = 1/(1+np.exp(-ld))
    return ld, prob


def gradient_check(x, y):
    p = init(76); _, g = loss_grad(x[:8], y[:8], p); err = 0.
    for name in p:
        for ix in list(np.ndindex(p[name].shape))[:8]:
            old = p[name][ix]; eps = 1e-5
            p[name][ix] = old+eps; plus = loss_grad(x[:8], y[:8], p)[0]
            p[name][ix] = old-eps; minus = loss_grad(x[:8], y[:8], p)[0]
            p[name][ix] = old
            err = max(err, abs((plus-minus)/(2*eps)-g[name][ix]))
    assert err < 1e-7, err
    return err


def train(x, y, dev_x, dev_y, seed):
    p = init(seed); m = {k:np.zeros_like(v) for k,v in p.items()}
    v = {k:np.zeros_like(a) for k,a in p.items()}; curve=[]
    for t in range(1, 1001):
        loss, g = loss_grad(x, y, p)
        for k in p:
            m[k] = .9*m[k]+.1*g[k]; v[k] = .999*v[k]+.001*g[k]**2
            p[k] -= .01*(m[k]/(1-.9**t))/(np.sqrt(v[k]/(1-.999**t))+1e-8)
        if t in [1, 100, 300, 600, 1000]:
            curve.append(dict(step=t, train_nll=loss_grad(x,y,p)[0],
                              dev_nll=loss_grad(dev_x,dev_y,p)[0],
                              dev_accuracy=float(np.mean(forward(dev_x,p)[0].argmax(1)==dev_y))))
    return p, curve


def choose_top4(dev_x, dev_y, p, layer):
    zc, hc = forward(dev_x, p); zb, hb = forward(corrupt(dev_x), p)
    base = scores(zb, dev_y)[0]; effects=[]
    for j in range(16):
        h = hb[layer-1].copy(); h[:,j] = hc[layer-1][:,j]
        effects.append(float(np.mean(scores(resume(h,layer,p),dev_y)[0]-base)))
    return np.argsort(effects)[-4:].tolist()


def evaluate(x, y, p, rotation, top4):
    records=[]; zc, hc = forward(x,p); ldc, pc = scores(zc,y)
    for condition in ['selected_flip','unselected_flip']:
        xb = corrupt(x, condition=='selected_flip'); zb, hb = forward(xb,p)
        ldb, pb = scores(zb,y); denom = ldc-ldb
        valid = (condition=='selected_flip') & (denom > 1e-6)
        for layer in [1,2]:
            for basis in (['native','rotated'] if layer==2 else ['native']):
                q = rotation if basis=='rotated' else np.eye(16)
                clean = hc[layer-1]@q; bad = hb[layer-1]@q
                cases = [('none', [], 'clean'), ('self', list(range(16)), 'self'),
                         ('full', list(range(16)), 'clean')]
                cases += [('single_'+str(j), [j], 'clean') for j in range(16)]
                if basis=='native':
                    cases += [('top4',top4[layer],'clean'), ('reverse_top4',top4[layer],'reverse'),
                              ('zero_top4',top4[layer],'zero'), ('shuffled_top4',top4[layer],'shuffle')]
                for name, indices, source in cases:
                    h = clean.copy() if source=='reverse' else bad.copy()
                    donor = {'clean':clean,'self':bad,'reverse':bad,
                             'zero':np.zeros_like(bad),'shuffle':np.roll(clean,1,axis=0)}[source]
                    h[:,indices] = donor[:,indices]
                    z = resume(h@q.T,layer,p); ld, prob = scores(z,y)
                    if name in ['none','self']: assert np.allclose(z,zb,atol=1e-12)
                    if name=='full': assert np.allclose(z,zc,atol=1e-12)
                    for i in range(len(x)):
                        records.append(dict(condition=condition,layer=layer,basis=basis,
                            intervention=name,indices=indices,source=source,sample=i,
                            logits=z[i].tolist(),ld=float(ld[i]),prob=float(prob[i]),
                            clean_ld=float(ldc[i]),corrupt_ld=float(ldb[i]),
                            delta_ld=float(ld[i]-ldb[i]),delta_prob=float(prob[i]-pb[i]),
                            recovery=float((ld[i]-ldb[i])/denom[i]) if valid[i] else None,
                            degradation=float(ldc[i]-ld[i]) if source=='reverse' else None))
    return records


def summarize(records):
    groups={}
    for row in records:
        key=(row['condition'],row['layer'],row['basis'],row['intervention'])
        groups.setdefault(key,[]).append(row)
    output=[]
    for key, rows in groups.items():
        rec=[r['recovery'] for r in rows if r['recovery'] is not None]
        output.append(dict(zip(['condition','layer','basis','intervention'],key)) | dict(
            n=len(rows),accuracy=float(np.mean([r['ld']>0 for r in rows])),
            mean_delta_ld=float(np.mean([r['delta_ld'] for r in rows])),
            mean_delta_prob=float(np.mean([r['delta_prob'] for r in rows])),
            mean_recovery=float(np.mean(rec)) if rec else None,
            valid_recovery_n=len(rec)))
    return output


def main():
    ap=argparse.ArgumentParser();ap.add_argument('--out',type=Path,default=Path(__file__).parent/'results')
    args=ap.parse_args();args.out.mkdir(parents=True,exist_ok=True)
    rng=np.random.default_rng(760)
    x,y=dataset(rng,64); dx,dy=dataset(rng,16); tx,ty=dataset(rng,8)
    assert not (set(map(tuple,x)) & set(map(tuple,dx)))
    assert not (set(map(tuple,tx)) & (set(map(tuple,x))|set(map(tuple,dx))))
    np.savez(args.out/'data.npz',train_x=x,train_y=y,dev_x=dx,dev_y=dy,eval_x=tx,eval_y=ty)
    grad_err=gradient_check(x,y)
    summary=dict(python=platform.python_version(),numpy=np.__version__,dtype='float64',
                 parameters=402,gradient_max_abs_error=grad_err,steps=1000,lr=.01,
                 data_seed=760,runs=[])
    print('shapes: train',x.shape,'dev',dx.shape,'eval',tx.shape,'hidden',(64,16),'logits',(64,2))
    print('finite-difference gradient max error:',grad_err)
    for seed in [76,77,78]:
        p,curve=train(x,y,dx,dy,seed)
        q,_=np.linalg.qr(np.random.default_rng(seed+1000).normal(size=(16,16)))
        assert np.allclose(q@q.T,np.eye(16),atol=1e-12)
        top4={layer:choose_top4(dx,dy,p,layer) for layer in [1,2]}
        np.savez(args.out/f'model_{seed}.npz',**p,rotation=q)
        frozen={k:v.copy() for k,v in p.items()}
        rows=evaluate(tx,ty,p,q,top4)
        assert all(np.array_equal(p[k],v) for k,v in frozen.items())
        (args.out/f'records_{seed}.json').write_text(json.dumps(rows,indent=2,allow_nan=False)+'\n')
        run=dict(seed=seed,curve=curve,top4=top4,clean_accuracy=float(np.mean(forward(tx,p)[0].argmax(1)==ty)),
                 groups=summarize(rows))
        summary['runs'].append(run)
        print('seed',seed,'clean accuracy',run['clean_accuracy'],'top4',top4,'records',len(rows))
        for g in run['groups']:
            if g['condition']=='selected_flip' and g['basis']=='native' and g['intervention'] in ['none','full','top4']:
                print(g)
    (args.out/'summary.json').write_text(json.dumps(summary,indent=2,allow_nan=False)+'\n')
    print('PASS: full/self/none controls, orthogonal basis, frozen weights and finite metrics')


if __name__=='__main__': main()
