"""NumPy causal RNN protocol experiment; no pretrained LLM or paper benchmark.
Run with Python >=3.9 and numpy. Saves data, checkpoints, traces, predictions.
"""
import itertools
import json
import platform
from pathlib import Path
import numpy as np

OUT = Path(__file__).resolve().parent
TOKENS = ['PAD', 'BOS', 'EOS', 'USER', 'ASSISTANT', 'TASK', 'INPUT', 'ANSWER', 'COPY', 'FLIP', '0', '1']
V = len(TOKENS)
ID = {s: i for i, s in enumerate(TOKENS)}
H = 48


def row(bits, task, template='A'):
    prompt = (['BOS', 'USER', task] + list(bits) + ['ASSISTANT'] if template == 'A'
              else ['BOS', 'TASK', task, 'INPUT'] + list(bits) + ['ANSWER'])
    answer = list(bits if task == 'COPY' else ''.join('1' if x == '0' else '0' for x in bits)) + ['EOS']
    return {'group': bits, 'task': task, 'template': template, 'prompt': prompt, 'answer': answer}


def batch(rows, full=False):
    length = max(len(r['prompt']) + len(r['answer']) for r in rows) - 1
    x = np.zeros((len(rows), length), dtype=int)
    y = x.copy()
    mask = np.zeros_like(x, dtype=float)
    for b, r in enumerate(rows):
        ids = [ID[t] for t in r['prompt'] + r['answer']]
        x[b, :len(ids)-1], y[b, :len(ids)-1] = ids[:-1], ids[1:]
        start = 0 if full else len(r['prompt']) - 1
        mask[b, start:len(ids)-1] = 1
    assert mask.sum() > 0
    return x, y, mask


def init(seed):
    rng = np.random.default_rng(seed)
    return {'E': rng.normal(0, .12, (V, H)), 'W': rng.normal(0, .12, (H, H)),
            'b': np.zeros(H), 'U': rng.normal(0, .12, (H, V)), 'c': np.zeros(V)}


def loss_grad(p, rows, full=False):
    x, y, mask = batch(rows, full)
    b, t = x.shape
    hs = [np.zeros((b, H))]
    probs = []
    loss = 0.
    for j in range(t):
        h = np.tanh(p['E'][x[:, j]] + hs[-1] @ p['W'] + p['b'])
        hs.append(h)
        z = h @ p['U'] + p['c']
        z -= z.max(axis=1, keepdims=True)
        logp = z - np.log(np.exp(z).sum(axis=1, keepdims=True))
        probs.append(np.exp(logp))
        loss -= (logp[np.arange(b), y[:, j]] * mask[:, j]).sum() / mask.sum()
    g = {k: np.zeros_like(a) for k, a in p.items()}
    dh = np.zeros((b, H))
    for j in reversed(range(t)):
        dz = probs[j].copy()
        dz[np.arange(b), y[:, j]] -= 1
        dz *= mask[:, j, None] / mask.sum()
        g['U'] += hs[j+1].T @ dz
        g['c'] += dz.sum(axis=0)
        da = (dh + dz @ p['U'].T) * (1 - hs[j+1]**2)
        np.add.at(g['E'], x[:, j], da)
        g['W'] += hs[j].T @ da
        g['b'] += da.sum(axis=0)
        dh = da @ p['W'].T
    return float(loss), g


class Adam:
    def __init__(self, p):
        self.m = {k: np.zeros_like(v) for k, v in p.items()}
        self.v = {k: np.zeros_like(v) for k, v in p.items()}
        self.t = 0

    def step(self, p, g, lr=.003):
        self.t += 1
        norm = np.sqrt(sum(float((v*v).sum()) for v in g.values()))
        for k in p:
            a = g[k] / max(1., norm)
            self.m[k] = .9*self.m[k] + .1*a
            self.v[k] = .999*self.v[k] + .001*a*a
            p[k] -= lr * (self.m[k]/(1-.9**self.t)) / (np.sqrt(self.v[k]/(1-.999**self.t))+1e-8)


def generate(p, r):
    h = np.zeros(H)
    for token in r['prompt']:
        h = np.tanh(p['E'][ID[token]] + h @ p['W'] + p['b'])
    answer = []
    for _ in range(8):
        token = int(np.argmax(h @ p['U'] + p['c']))
        answer.append(TOKENS[token])
        if token == ID['EOS']:
            break
        h = np.tanh(p['E'][token] + h @ p['W'] + p['b'])
    return answer


def evaluate(p, rows):
    predictions = [{**r, 'generated': generate(p, r)} for r in rows]
    exact = [r['generated'] == r['answer'] for r in predictions]
    return {'nll': loss_grad(p, rows)[0], 'exact_match': float(np.mean(exact)),
            'per_task': {task: float(np.mean([e for e, r in zip(exact, rows) if r['task'] == task]))
                         for task in ['COPY', 'FLIP'] if any(r['task'] == task for r in rows)},
            'predictions': predictions}


def main():
    groups = [''.join(x) for x in itertools.product('01', repeat=4)]
    np.random.default_rng(670).shuffle(groups)
    split = {'train': groups[:8], 'dev': groups[8:12], 'test': groups[12:]}
    assert all(not set(split[a]) & set(split[b]) for a,b in [('train','dev'),('train','test'),('dev','test')])
    data = {s: [row(g, task) for g in gs for task in ['COPY','FLIP']] for s,gs in split.items()}
    pretrain = [row(''.join(bits), task) for bits in itertools.product('01', repeat=3) for task in ['COPY','FLIP']]
    (OUT/'data.json').write_text(json.dumps({'split': split, **data, 'pretrain': pretrain}, indent=2))
    # Finite difference checks all parameter families, including a masked-prefix embedding.
    p = init(67)
    loss, grad = loss_grad(p, data['train'][:2])
    errors = []
    for key, idx in [('E',(ID['USER'],0)),('W',(0,0)),('b',(0,)),('U',(0,0)),('c',(0,))]:
        old = p[key][idx]
        p[key][idx] = old + 1e-5
        plus = loss_grad(p, data['train'][:2])[0]
        p[key][idx] = old - 1e-5
        minus = loss_grad(p, data['train'][:2])[0]
        p[key][idx] = old
        errors.append(abs((plus-minus)/2e-5 - grad[key][idx]))
    assert max(errors) < 1e-6
    x,y,m = batch(data['train'])
    assert np.all(m.sum(axis=1) == 5)
    assert y[0,np.flatnonzero(m[0])[0]] == ID[data['train'][0]['answer'][0]]
    # Duplicating a batch must preserve the token-mean loss.
    assert abs(loss_grad(p,[data['train'][0]])[0] - loss_grad(p,[data['train'][0]]*2)[0]) < 1e-12
    result = {'python': platform.python_version(), 'numpy': np.__version__, 'seed_data':670,
              'shapes': {'input':list(x.shape),'logits':[len(x),x.shape[1],V]},
              'parameters':sum(v.size for v in p.values()), 'gradient_max_error':max(errors),
              'config':{'pretrain_steps':600,'sft_steps':400,'batch':16,'lr':.003,'checkpoints':[20,100,200,400]}, 'runs':[]}
    for seed in [67,68,69]:
        p = init(seed)
        opt = Adam(p)
        for _ in range(600):
            opt.step(p, loss_grad(p, pretrain, full=True)[1])
        base = {k:v.copy() for k,v in p.items()}
        np.savez(OUT/f'base_{seed}.npz', **base)
        base_eval = evaluate(base, data['test'])
        for variant in ['copy_only','balanced','full_loss']:
            p = {k:v.copy() for k,v in base.items()}
            rows = ([r for r in data['train'] if r['task']=='COPY']*2 if variant=='copy_only' else data['train'])
            opt = Adam(p)
            trace, best, best_score, best_step = [], None, float('inf'), None
            for step in range(1,401):
                opt.step(p,loss_grad(p,rows,full=variant=='full_loss')[1])
                if step in [20,100,200,400]:
                    dev = loss_grad(p,data['dev'])[0]
                    trace.append({'step':step,'train_response_nll':loss_grad(p,rows)[0],'dev_response_nll':dev})
                    if dev < best_score:
                        best, best_score, best_step = {k:v.copy() for k,v in p.items()}, dev, step
            np.savez(OUT/f'{variant}_{seed}.npz', **best)
            run = {'seed':seed,'variant':variant,'selected_step':best_step,'trace':trace,
                   'supervised_tokens_per_step':int(batch(rows,full=variant=='full_loss')[2].sum()),
                   'train':evaluate(best,rows),'test':evaluate(best,data['test']),
                   'template_shift':evaluate(best,[row(r['group'],r['task'],'B') for r in data['test']]),
                   'base_test':base_eval}
            result['runs'].append(run)
            print(seed,variant,'step',best_step,'test',run['test']['exact_match'],'shift',run['template_shift']['exact_match'],flush=True)
    result['summary'] = {}
    for variant in ['copy_only','balanced','full_loss']:
        rs=[r for r in result['runs'] if r['variant']==variant]
        result['summary'][variant]={k:{'mean':float(np.mean([r[k]['exact_match'] for r in rs])),
                                                  'sample_sd':float(np.std([r[k]['exact_match'] for r in rs],ddof=1))}
                                     for k in ['train','test','template_shift','base_test']}
    (OUT/'results.json').write_text(json.dumps(result,indent=2))
    print(json.dumps({k:v for k,v in result.items() if k!='runs'},indent=2))


if __name__ == '__main__':
    main()
