"""Reproduce the recorded NN/PINN and held-out graph learning demonstrations.
Python 3 + NumPy only. No test labels enter optimization or model selection.
Run: OPENBLAS_NUM_THREADS=1 python3 scripts/generate-learning-replays.py
"""
import argparse, base64, hashlib, json, time
from pathlib import Path
import numpy as np

ROOT=Path(__file__).resolve().parents[1]
OUT=ROOT/'public/learning-lab'
SEED=11

def score(pred,y,mask=None):
    if mask is not None: pred,y=pred[mask],y[mask]
    err=np.square(pred-y)
    # Variance-weighted R² over output channels (channel-specific means).
    den=np.square(y-y.mean(axis=0)).sum()
    return {'r2':float(1-err.sum()/den) if den>1e-15 else None,'rmse':float(np.sqrt(err.mean()))}

def tensor(a):
    a=np.asarray(a,dtype='<f4')
    return {'shape':list(a.shape),'dtype':'<f4','base64':base64.b64encode(a.tobytes()).decode()}

def rounded(a): return np.asarray(a).round(10).tolist()

class Net:
    def __init__(self,kind,inputs=4,outputs=2,width=32,seed=SEED):
        self.kind=kind; self.p={}; self.g={}; self.m={}; self.v={}; self.t=0
        rng=np.random.default_rng(seed)
        def layer(name,i,o):
            self.p[name+'w']=rng.normal(size=(i,o))*np.sqrt(2/(i+o))
            self.p[name+'b']=np.zeros(o)
        if kind=='linear': layer('a',9,2)
        else:
            layer('e',inputs,width)
            layer('h',width,width)
            if kind=='gnn': layer('n',width,width)
            layer('o',width,outputs)
            if kind in ['residual','gnn']: layer('a',9,2)
        self.zero()
    def zero(self): self.g={k:np.zeros_like(v) for k,v in self.p.items()}
    def count(self): return sum(v.size for v in self.p.values())
    def forward(self,x,a=None,f=None):
        p=self.p
        if self.kind=='linear': return f@p['aw']+p['ab'],(x,a,f,None,None,None)
        h=np.tanh(x@p['ew']+p['eb']); agg=None
        z=h@p['hw']+p['hb']
        if self.kind=='gnn': agg=np.einsum('ij,bjk->bik',a,h);z+=agg@p['nw']+p['nb']
        h2=np.tanh(z);y=h2@p['ow']+p['ob']
        if self.kind in ['residual','gnn']:y+=f@p['aw']+p['ab']
        return y,(x,a,f,h,agg,h2)
    def backward(self,d,c):
        x,a,f,h,agg,h2=c;p=self.p;g=self.g
        def add(name,u,d):
            g[name+'w']+=u.reshape(-1,u.shape[-1]).T@d.reshape(-1,d.shape[-1]);g[name+'b']+=d.reshape(-1,d.shape[-1]).sum(0)
        if self.kind in ['linear','residual','gnn']:add('a',f,d)
        if self.kind=='linear':return
        add('o',h2,d);dz=(d@p['ow'].T)*(1-h2*h2);add('h',h,dz);dh=dz@p['hw'].T
        if self.kind=='gnn':
            add('n',agg,dz);dh+=np.einsum('ji,bjk->bik',a,dz@p['nw'].T)
        de=dh*(1-h*h);add('e',x,de)
    def step(self,lr=.002):
        self.t+=1
        for k,p in self.p.items():
            g=self.g[k]
            self.m[k]=.9*self.m.get(k,np.zeros_like(g))+.1*g
            self.v[k]=.999*self.v.get(k,np.zeros_like(g))+.001*g*g
            p-=lr*(self.m[k]/(1-.9**self.t))/(np.sqrt(self.v[k]/(1-.999**self.t))+1e-8)

def pinn():
    x=np.linspace(0,1,161)[:,None];xt=np.linspace(0,.35,8)[:,None];xc=np.linspace(.001,.999,64)[:,None]
    delta=.001
    cases=[]
    for case,title,fn,df in [
      ('hardening','Nonlinear strain hardening',lambda z:.6*z+.4*z**3,lambda z:.6+1.2*z*z),
      ('saturation','Saturating stress response',lambda z:1-np.exp(-3*z),lambda z:3*np.exp(-3*z))]:
        truth=fn(x);yt=fn(xt);test=(x[:,0]>.35+1e-10)
        models=[Net('mlp',inputs=1,outputs=1,width=24),Net('mlp',inputs=1,outputs=1,width=24)]
        frames=[]
        for step in range(2401):
            if step%40==0:
                preds=[m.forward(x)[0] for m in models]
                frames.append({'step':step,'predictions':[rounded(v[:,0]) for v in preds],
                  'scores':[score(v[test],truth[test]) for v in preds],
                  'train_scores':[score(m.forward(xt)[0],yt) for m in models]})
            if step==2400:break
            for j,m in enumerate(models):
                m.zero();yp,c=m.forward(xt);d=2*(yp-yt)/len(xt);m.backward(d,c)
                if j==1:
                    plus,cp=m.forward(xc+delta);minus,cm=m.forward(xc-delta)
                    residual=(plus-minus)/(2*delta)-df(xc)
                    dr=2*.15*residual/len(xc)/(2*delta)
                    m.backward(dr,cp);m.backward(-dr,cm)
                m.step(.003)
        cases.append({'id':case,'title':title,'truth':rounded(truth[:,0]),'training_x':rounded(xt[:,0]),'training_y':rounded(yt[:,0]),'frames':frames,
          'law':'σ = 0.6ε + 0.4ε³; dσ/dε = 0.6 + 1.2ε²' if case=='hardening' else 'σ = 1 − exp(−3ε); dσ/dε = 3 exp(−3ε)'})
        print('PINN',case,'final',frames[-1]['scores'],flush=True)
    data={'schema_version':1,'kind':'executed_synthetic_constitutive_training','seed':SEED,'x':rounded(x[:,0]),
      'models':[{'id':'nn','name':'NN · data only','parameters':models[0].count()},{'id':'pinn','name':'PINN · data + constitutive law','parameters':models[1].count()}],
      'protocol':{'steps':2400,'optimizer':'Adam','learning_rate':.003,'widths':[1,24,24,1],'activation':'tanh','training_points':8,'collocation_points':64,'physics_weight':.15,'derivative':'central difference, h=0.001','test_region':'ε > 0.35; never used as observed labels during training','units':'dimensionless normalized strain and stress','same_initial_weights':True},'cases':cases}
    (OUT/'pinn-replay.json').write_text(json.dumps(data,separators=(',',':'),allow_nan=False))


def load_graphs(folder):
    datasets={}
    for id in ['Exp3','Exp4','Exp5','Probe1_22','Probe2_7']:
        path=folder/f'{id}.json';raw=path.read_bytes();d=json.loads(raw);g=d['graph'];o=d['observations'];n=len(g['x0'])
        def dec(t):return np.frombuffer(base64.b64decode(t['base64']),dtype=t['dtype']).reshape(t['shape']).astype(float)
        xy=np.array(g['x0']);xy-=xy.mean(0)
        phase=np.array(o['phases']);x=np.empty((len(phase),n,4));x[:,:,:2]=xy;x[:,:,2]=phase[:,None];x[:,:,3]=phase[:,None]**2
        f=np.stack([np.ones(x.shape[:2]),x[:,:,0],x[:,:,1],x[:,:,2],x[:,:,2]*x[:,:,0],x[:,:,2]*x[:,:,1],x[:,:,3],x[:,:,3]*x[:,:,0],x[:,:,3]*x[:,:,1]],-1)
        a=np.zeros((n,n));src,dst=g['edge_index'];a[dst,src]=1;a/=np.maximum(1,a.sum(1))[:,None]
        datasets[id]={'x':x,'f':f,'a':a,'y':dec(o['displacements']),'mask':dec(o['mask']).astype(bool),'raw':d,'sha256':hashlib.sha256(raw).hexdigest()}
    return datasets

def gnn(folder):
    ds=load_graphs(folder);train=[ds[k] for k in ['Exp3','Exp4','Exp5']];valid=ds['Probe2_7'];test=ds['Probe1_22']
    models=[Net(k) for k in ['linear','mlp','residual','gnn']]
    rng=np.random.default_rng(SEED);frames=[];selected=np.linspace(0,len(test['x'])-1,13).round().astype(int)
    best=[(float('inf'),0) for m in models]
    for step in range(1801):
        if step%30==0:
            scores=[];vs=[];preds=[];ts=[]
            for j,m in enumerate(models):
                yp=m.forward(test['x'],test['a'],test['f'])[0];scores.append(score(yp,test['y'],test['mask']));preds.append(tensor(yp[selected]))
                vp=m.forward(valid['x'],valid['a'],valid['f'])[0];v=score(vp,valid['y'],valid['mask']);vs.append(v)
                if v['rmse']<best[j][0]:best[j]=(v['rmse'],step)
                ty=[];py=[]
                for d in train:
                    y=m.forward(d['x'],d['a'],d['f'])[0];ty.append(d['y'][d['mask']]);py.append(y[d['mask']])
                ts.append(score(np.concatenate(py),np.concatenate(ty)))
            frames.append({'step':step,'scores':scores,'validation_scores':vs,'train_scores':ts,'predictions':preds})
            if step%300==0:print('GNN',step,[round(s['r2'],4) for s in scores],flush=True)
        if step==1800:break
        batches=[rng.choice(len(d['x']),size=min(12,len(d['x'])),replace=False) for d in train]
        total=sum(d['mask'][b].sum()*2 for d,b in zip(train,batches))
        for m in models:
            m.zero()
            for d,b in zip(train,batches):
                yp,c=m.forward(d['x'][b],d['a'],d['f'][b]);dy=2*(yp-d['y'][b])*d['mask'][b,:,None]/total;m.backward(dy,c)
            m.step(.002)
    names=['Polynomial regression','MLP','Residual MLP','GNN · mean messages']
    data={'schema_version':1,'kind':'executed_five_specimen_controlled_training','seed':SEED,
      'models':[{'id':m.kind,'name':names[i],'parameters':m.count(),'best_validation_step':best[i][1]} for i,m in enumerate(models)],
      'protocol':{'steps':1800,'optimizer':'Adam','learning_rate':.002,'batch_frames_per_specimen':12,'training_specimens':['Exp3_D12','Exp4_D14','Exp5_ID03'],'validation_specimens':['Probe2_7'],'test_specimens':['Probe1_22'],'normalization':'displacement / initial graph height','score':'variance-weighted R² of dx and dy on all valid held-out nodes and recorded phases','selection':'full fixed training budget; best recorded validation checkpoint marked separately','architecture_note':'MLP: 4→32→32→2 tanh. Residual MLP adds a 9-feature polynomial skip. GNN adds a learned mean-neighbor message transform to that residual model. Polynomial regression uses the same 9 explicit features. Parameter counts differ; this is not a capacity-matched research benchmark.','source_commit':test['raw']['source_commit'],'source_hashes':{k:v['sha256'] for k,v in ds.items()},'original_demo_preserved':True},
      'test':{'id':'Probe1_22','physical_id':test['raw']['physical_id'],'graph':test['raw']['graph'], 'frames':[test['raw']['observations']['frames'][i] for i in selected], 'phases':rounded(np.array(test['raw']['observations']['phases'])[selected]),'truth':tensor(test['y'][selected]),'mask':tensor(test['mask'][selected]),'all_valid_test_pairs':int(test['mask'].sum())},'frames':frames}
    (OUT/'gnn-replay.json').write_text(json.dumps(data,separators=(',',':'),allow_nan=False))
    print('GNN best validation steps',best,flush=True)

if __name__=='__main__':
    p=argparse.ArgumentParser();p.add_argument('--gnn-data',type=Path,default=ROOT/'public/gnn-cb/data');p.add_argument('--only',choices=['pinn','gnn','all'],default='all');args=p.parse_args();OUT.mkdir(parents=True,exist_ok=True)
    if args.only in ('pinn','all'):pinn()
    if args.only in ('gnn','all'):gnn(args.gnn_data)
