import json, numpy as np

M = json.load(open('model.json'))
P = {k: np.array(v, dtype=np.float64) for k, v in M['params'].items()}
vocab = M['vocab']; T, D, H, DH, F = M['T'], M['D'], M['H'], M['DH'], M['F']
stoi = {w: i for i, w in enumerate(vocab)}

def ln(x, g, b, eps=1e-5):
    mu = x.mean(-1, keepdims=True); var = ((x - mu) ** 2).mean(-1, keepdims=True)
    return (x - mu) / np.sqrt(var + eps) * g + b
def gelu(x):
    return 0.5 * x * (1 + np.tanh(np.sqrt(2 / np.pi) * (x + 0.044715 * x ** 3)))
def softmax(x):
    x = x - x.max(-1, keepdims=True); e = np.exp(x); return e / e.sum(-1, keepdims=True)

def fwd(words):
    ids = [stoi[w] for w in words][-T:]
    n = len(ids)
    x0 = P['E'][ids] + P['Pos'][:n]
    a = ln(x0, P['g1'], P['b1'])
    q = a @ P['Wq']; k = a @ P['Wk']; v = a @ P['Wv']
    att = np.zeros((n, D))
    for h in range(H):
        s = slice(h * DH, (h + 1) * DH)
        S = q[:, s] @ k[:, s].T / np.sqrt(DH)
        S = np.where(np.triu(np.ones((n, n), bool), 1), -1e9, S)
        A = softmax(S)
        att += (A @ v[:, s]) @ P['Wo'][s, :]
    x1 = x0 + att
    m = ln(x1, P['g2'], P['b2'])
    x2 = x1 + gelu(m @ P['W1'] + P['c1']) @ P['W2'] + P['c2']
    return x2 @ P['U']

refs = {}
prompts = ['the cat', 'the dog sat on', 'the big dog chased the', 'a dog sat on a', 'the mouse ate the cheese . the']
# coverage windows: every vocabulary word and every position 0..T-1 gets exercised
cover = vocab + vocab[:T]
prompts += [' '.join(cover[i:i + T]) for i in range(0, len(vocab), T)]
for p in prompts:
    refs[p] = fwd(p.split())[-1].round(6).tolist()

spec = []
for h in range(H):
    s = slice(h * DH, (h + 1) * DH)
    qk = P['Wq'][:, s] @ P['Wk'][:, s].T      # D x D
    ov = P['Wv'][:, s] @ P['Wo'][s, :]        # D x D
    sq = np.linalg.svd(qk, compute_uv=False)[:DH]
    so = np.linalg.svd(ov, compute_uv=False)[:DH]
    spec.append({'qk': sq.round(5).tolist(), 'ov': so.round(5).tolist()})
M['svd'] = spec
M['refs'] = refs
json.dump(M, open('model_full.json', 'w'), separators=(',', ':'))
print('ok', len(json.dumps(M)) // 1024, 'KB')
txt = json.dumps(M, separators=(',', ':')).replace('],[', '],\n[').replace('],"', '],\n"')
open('model.js', 'w').write('// generated by build_data.py from train.py output; do not edit\nvar MODEL = ' + txt + ';\n')
for h in range(H):
    print('head', h, 'QK sv', spec[h]['qk'], '\n       OV sv', spec[h]['ov'])
