"""CPU linear probing on frozen locally trained checkpoints; not an NLP benchmark."""
import itertools
import json
import hashlib
import platform
from pathlib import Path
import numpy as np

ROOT = Path(__file__).resolve().parent
LAMBDAS = [0., .001, .01, .1]


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


def data(seed=770):
    r = np.random.default_rng(seed)
    out = {}
    for split, repeats in [('train',64), ('dev',32), ('test',64)]:
        core = np.repeat(np.array(list(itertools.product([-1.,1.], repeat=3))), repeats, 0)
        x = np.c_[core, r.normal(0,.3,(len(core),2))]
        selected = (x[:,2]>0).astype(int)
        out[split+'_x'] = x
        out[split+'_task'] = (x[np.arange(len(x)),selected]>0).astype(int)
        out[split+'_unused'] = (x[np.arange(len(x)),1-selected]>0).astype(int)
        # Independent example-level balanced shuffle; NOT word-type control task.
        out[split+'_shuffle'] = r.permutation(out[split+'_unused'])
    return out


def loss_grad(h,y,w,lam):
    z = h@w
    loss = np.mean(np.logaddexp(0,z)-y*z) + .5*lam*np.sum(w[:-1]**2)
    prob = 1/(1+np.exp(-np.clip(z,-700,700)))
    g = h.T@(prob-y)/len(y)
    g[:-1] += lam*w[:-1]
    return float(loss), g


def fit(h,y,lam):
    # Fixed-step full-batch gradient descent; convex logistic objective.
    w=np.zeros(h.shape[1])
    lr = 1/(.25*np.linalg.norm(h,2)**2/len(h)+lam)
    for _ in range(2000):
        _,g=loss_grad(h,y,w,lam); w-=lr*g
    return w, float(np.linalg.norm(loss_grad(h,y,w,lam)[1]))


def metrics(z,y):
    return dict(accuracy=float(np.mean((z>=0)==y)),
                nll=float(np.mean(np.logaddexp(0,z)-y*z)))


def main():
    results=ROOT/'results';results.mkdir(exist_ok=True)
    d=data();np.savez(results/'data.npz',**d)
    # Verify split rows are genuinely distinct, despite shared eight core patterns.
    rows=[set(map(tuple,d[s+'_x'])) for s in ['train','dev','test']]
    assert all(not rows[i]&rows[j] for i in range(3) for j in range(i))
    rg=np.random.default_rng(771);h=np.c_[rg.normal(size=(9,5)),np.ones(9)]
    y=rg.integers(0,2,9);w=rg.normal(size=6);_,g=loss_grad(h,y,w,.01)
    errors=[]
    for j in range(6):
        a=w.copy();b=w.copy();a[j]+=1e-5;b[j]-=1e-5
        errors.append(abs((loss_grad(h,y,a,.01)[0]-loss_grad(h,y,b,.01)[0])/2e-5-g[j]))
    assert max(errors)<1e-7
    records=[];behavior=[]; hashes={}
    for seed in [76,77,78]:
        path=ROOT/'checkpoints'/f'model_{seed}.npz'
        hashes[path.name]=hashlib.sha256(path.read_bytes()).hexdigest()
        p=dict(np.load(path));before={k:v.copy() for k,v in p.items()}
        cache={s:forward(d[s+'_x'],p) for s in ['train','dev','test']}
        for layer in range(3):
            tr=cache['train'][0][layer];mu=tr.mean(0);sd=tr.std(0);sd=np.maximum(sd,1e-8)
            hs={s:np.c_[(cache[s][0][layer]-mu)/sd,np.ones(len(d[s+'_x']))] for s in cache}
            for target in ['task','unused','shuffle']:
                candidates=[]
                for lam in LAMBDAS:
                    w,grad=fit(hs['train'],d['train_'+target],lam)
                    m=metrics(hs['dev']@w,d['dev_'+target])
                    candidates.append((m['nll'],lam,w,grad))
                chosen=min(range(len(candidates)),key=lambda i:candidates[i][0])
                _,lam,w,grad=candidates[chosen]
                pred={s:hs[s]@w for s in hs}
                name=f's{seed}_l{layer}_{target}'
                np.savez(results/(name+'.npz'),w=w,mean=mu,scale=sd,**{s+'_logit':z for s,z in pred.items()},
                         candidate_w=np.stack([c[2] for c in candidates]))
                records.append(dict(seed=seed,layer=layer,target=target,lambda_=lam,chosen_index=chosen,
                                    dev_candidate_nll=[c[0] for c in candidates],gradient_norm=grad,
                                    **{s:metrics(pred[s],d[s+'_'+target]) for s in pred},file=name+'.npz'))
        x=d['test_x'];original=cache['test'][1];flipped=x.copy();pos=(x[:,2]>0).astype(int)
        flipped[np.arange(len(x)),1-pos]*=-1
        cf=forward(flipped,p)[1]
        # Exact output-nullspace projection control in h2; preserves both logits.
        h2=cache['test'][0][2];_,sv,vt=np.linalg.svd(p['W3'].T,full_matrices=True)
        rank=int(np.sum(sv>1e-10));basis=vt[:rank].T;proj=h2@basis@basis.T
        zproj=proj@p['W3']+p['b3'];err=float(np.max(np.abs(zproj-original)))
        assert err<1e-10
        # Freeze probe trained on original h2 and evaluate after projection (no retraining).
        probe=np.load(results/f's{seed}_l2_unused.npz')
        pz=np.c_[(proj-probe['mean'])/probe['scale'],np.ones(len(x))]@probe['w']
        behavior.append(dict(seed=seed,original_accuracy=float(np.mean(original.argmax(1)==d['test_task'])),
                             unused_flip_answer_change=float(np.mean(original.argmax(1)!=cf.argmax(1))),
                             unused_flip_max_logit_change=float(np.max(np.abs(original-cf))),
                             projection_rank=rank,projection_max_logit_error=err,
                             projected_unused_probe_accuracy=metrics(pz,d['test_unused'])['accuracy']))
        np.savez(results/f'behavior_{seed}.npz',original=original,flipped=cf,projected=zproj,
                 projected_probe=pz,basis=basis,projected_h2=proj)
        assert all(np.array_equal(p[k],before[k]) for k in p)
    aggregate={}
    for layer in range(3):
        for target in ['task','unused','shuffle']:
            vals=[r['test']['accuracy'] for r in records if r['layer']==layer and r['target']==target]
            aggregate[f'l{layer}_{target}']=dict(mean=float(np.mean(vals)),sd=float(np.std(vals,ddof=1)))
    summary=dict(environment=dict(python=platform.python_version(),numpy=np.__version__,device='CPU',dtype='float64'),
                 data_seed=770,model_seeds=[76,77,78],shapes={s:list(d[s+'_x'].shape) for s in ['train','dev','test']},
                 lambdas=LAMBDAS,steps=2000,checkpoint_hashes=hashes,gradient_max_error=max(errors),
                 records=records,behavior=behavior,aggregate=aggregate)
    (results/'summary.json').write_text(json.dumps(summary,indent=2)+'\n')
    print(json.dumps({k:v for k,v in summary.items() if k!='records'},indent=2))
    print('PASS: finite gradients, distinct split rows, frozen weights, output-nullspace invariance')

if __name__=='__main__':main()
