"""Learn tiny dynamics from episodes; compare teacher forcing and open loop."""
import argparse
import json
import platform
from pathlib import Path
import numpy as np

SEEDS = (74, 75, 76)
MODES = ('position_only', 'linear_state', 'sine_state')
H = 40

def transition(s, a):
    x, v = s[..., 0], s[..., 1]
    vn = .98 * v - .1 * np.sin(x) + .08 * a
    return np.stack((x + .1 * vn, vn), -1)

def episodes(seed, count, ood=False):
    rng = np.random.default_rng(seed)
    s = np.zeros((count, H + 1, 2))
    s[:, 0, 0] = rng.uniform(1.5, 2.5, count) if ood else rng.uniform(-.5, .5, count)
    s[:, 0, 1] = rng.uniform(-.3, .3, count)
    a = rng.choice(np.array([-.5, 0., .5]), (count, H))
    for t in range(H):
        s[:, t+1] = transition(s[:, t], a[:, t])
    return {'states': s, 'actions': a}

def features(s, a, mode):
    x, v = s[..., 0], s[..., 1]
    columns = [x, a, np.ones_like(x)] if mode == 'position_only' else [x, v, a, np.ones_like(x)]
    if mode == 'sine_state':
        columns.append(np.sin(x))
    return np.stack(columns, -1)

def predict(s, a, w, mode):
    return features(s, a, mode) @ w

def rollout(data, w, mode, correction=0):
    truth, a = data['states'], data['actions']
    # Future truth is read only in explicitly labelled correction condition.
    p = np.empty_like(truth); p[:, 0] = truth[:, 0]
    for t in range(H):
        inp = truth[:, t] if correction and t % correction == 0 else p[:, t]
        p[:, t+1] = predict(inp, a[:, t], w, mode)
    return p

def score(truth, pred):
    sq = (truth[:, 1:] - pred[:, 1:]) ** 2
    return {'position_mse_by_horizon': sq[..., 0].mean(0).tolist(),
            'velocity_mse_by_horizon': sq[..., 1].mean(0).tolist(),
            'position_mse_all': float(sq[..., 0].mean())}

def main(out):
    out.mkdir(parents=True, exist_ok=True)
    runs=[]
    for seed in SEEDS:
        data = {k: episodes(seed*100+i, n, k=='ood') for i, (k,n) in enumerate([('train',128),('dev',32),('id',64),('ood',64)])}
        np.savez_compressed(out/f'data_{seed}.npz', **{f'{k}_{v}': x for k,d in data.items() for v,x in d.items()})
        train=data['train']; n=128*H
        weights={}; predictions={}
        for mode in MODES:
            x=features(train['states'][:,:-1],train['actions'],mode).reshape(n,-1)
            y=train['states'][:,1:].reshape(n,2)
            # Fixed full-rank least squares: no random initialization or tuning.
            w, _, rank, _ = np.linalg.lstsq(x,y,rcond=None)
            assert rank==x.shape[1]
            assert np.max(np.abs(x.T @ (x@w-y))/n)<1e-10
            weights[mode]=w
            row={'seed':seed,'mode':mode,'parameters':int(w.size),'rank':int(rank),'metrics':{}}
            print('FIT',seed,mode,'X',x.shape,'Y',y.shape,'W',w.shape)
            for split in ('dev','id','ood'):
                d=data[split];truth=d['states']
                tf=predict(truth[:,:-1],d['actions'],w,mode)
                p=rollout(d,w,mode); corrected=rollout(d,w,mode,5)
                assert np.allclose(p[:,1],tf[:,0])
                poisoned={'states':truth.copy(),'actions':d['actions']}
                poisoned['states'][:,1:]=999
                assert np.array_equal(p,rollout(poisoned,w,mode))
                assert np.isfinite(p).all()
                row['metrics'][split]={'teacher_position_mse':float(((tf-truth[:,1:])[...,0]**2).mean()),
                     'teacher_velocity_mse':float(((tf-truth[:,1:])[...,1]**2).mean()),
                     'open':score(truth,p),'correct5':score(truth,corrected)}
                predictions[f'{mode}_{split}_teacher']=tf
                predictions[f'{mode}_{split}_open']=p
                predictions[f'{mode}_{split}_correct5']=corrected
            runs.append(row)
        np.savez_compressed(out/f'weights_{seed}.npz',**weights)
        np.savez_compressed(out/f'predictions_{seed}.npz',**predictions)
    aggregate={}
    for mode in MODES:
        aggregate[mode]={}
        for split in ('id','ood'):
            rows=[r['metrics'][split] for r in runs if r['mode']==mode]
            vals={'teacher_x':[r['teacher_position_mse'] for r in rows],
                  'open_h1':[r['open']['position_mse_by_horizon'][0] for r in rows],
                  'open_h10':[r['open']['position_mse_by_horizon'][9] for r in rows],
                  'open_h20':[r['open']['position_mse_by_horizon'][19] for r in rows],
                  'open_h40':[r['open']['position_mse_by_horizon'][39] for r in rows],
                  'open_all':[r['open']['position_mse_all'] for r in rows],
                  'correct5_all':[r['correct5']['position_mse_all'] for r in rows]}
            aggregate[mode][split]={k:{'mean':float(np.mean(v)),'seed_sd':float(np.std(v,ddof=1))} for k,v in vals.items()}
    # Same position/action with opposite velocities: non-Markov observation fixture.
    s=np.array([[0.,.3],[0.,-.3]]); a=np.zeros(2)
    nxt=transition(s,a)
    assert np.isclose(nxt[0,0]-nxt[1,0],.0588)
    summary={'python':platform.python_version(),'numpy':np.__version__,'seed_role':'independent episode generation, deterministic least squares',
             'seeds':SEEDS,'horizon':H,'split_counts':{'train':128,'dev':32,'id':64,'ood':64},
             'config':{'dt':.1,'damping':.98,'force':.1,'action_gain':.08,'actions':[-.5,0,.5]},
             'aliasing_next_positions':nxt[:,0].tolist(),'runs':runs,'aggregate':aggregate}
    (out/'summary.json').write_text(json.dumps(summary,indent=2)+'\n')
    print(json.dumps(aggregate,indent=2));print('PASS shapes/rank/normal equations/h1/no future truth/finite/aliasing')

if __name__=='__main__':
    parser=argparse.ArgumentParser();parser.add_argument('--out',type=Path,default=Path(__file__).parent/'results')
    main(parser.parse_args().out)
