"""
One attention head learns a key-value lookup from scratch.
Companion to https://vivekchalla.dev/blog/how-llms-work-inside/

Context: 6 tokens, each a (letter, digit) pair, e.g. "B7".  Final token: "B?".
Target: the digit paired with the queried letter.  Embeddings are FIXED random
vectors (letter vector + digit vector); only W_Q, W_K, W_V, W_O are trained, with
hand-written gradients (checked against finite differences).
"""
import json
import numpy as np

rng = np.random.default_rng(7)
LETTERS = list("ABCDEFGH")
DIGITS = list("0123456789")
D, DK, N = 16, 8, 6  # model dim, head dim, context pairs

E_let = rng.normal(0, 1, (len(LETTERS), D)) / np.sqrt(D)
E_dig = rng.normal(0, 1, (len(DIGITS) + 1, D)) / np.sqrt(D)  # last row = "?"
QMARK = len(DIGITS)


def make_batch(B, rng):
    lets = np.stack([rng.permutation(len(LETTERS))[:N] for _ in range(B)])
    digs = rng.integers(0, len(DIGITS), (B, N))
    pick = rng.integers(0, N, B)
    qlet = lets[np.arange(B), pick]
    y = digs[np.arange(B), pick]
    X = np.concatenate([E_let[lets] + E_dig[digs],
                        (E_let[qlet] + E_dig[QMARK])[:, None, :]], axis=1)  # B,N+1,D
    return X, y, lets, digs, qlet


def init(rng):
    s = 1 / np.sqrt(D)
    return {"WQ": rng.normal(0, s, (D, DK)), "WK": rng.normal(0, s, (D, DK)),
            "WV": rng.normal(0, s, (D, DK)), "WO": rng.normal(0, 1 / np.sqrt(DK), (DK, len(DIGITS)))}


def forward(P, X, y):
    xq = X[:, -1, :]                       # the "?" token asks the question
    q = xq @ P["WQ"]                       # B,DK
    K = X @ P["WK"]                        # B,T,DK
    V = X @ P["WV"]                        # B,T,DK
    s = np.einsum("bd,btd->bt", q, K) / np.sqrt(DK)
    a = np.exp(s - s.max(1, keepdims=True)); a /= a.sum(1, keepdims=True)
    o = np.einsum("bt,btd->bd", a, V)
    z = o @ P["WO"]
    p = np.exp(z - z.max(1, keepdims=True)); p /= p.sum(1, keepdims=True)
    loss = -np.log(p[np.arange(len(y)), y] + 1e-12).mean()
    return loss, dict(xq=xq, q=q, K=K, V=V, a=a, o=o, p=p)


def backward(P, X, y, c):
    B = len(y)
    dz = c["p"].copy(); dz[np.arange(B), y] -= 1; dz /= B
    g = {"WO": c["o"].T @ dz}
    do = dz @ P["WO"].T                                   # B,DK
    dV = c["a"][:, :, None] * do[:, None, :]              # B,T,DK
    da = np.einsum("btd,bd->bt", c["V"], do)
    ds = c["a"] * (da - (c["a"] * da).sum(1, keepdims=True))
    dq = np.einsum("bt,btd->bd", ds, c["K"]) / np.sqrt(DK)
    dK = ds[:, :, None] * c["q"][:, None, :] / np.sqrt(DK)
    g["WQ"] = c["xq"].T @ dq
    g["WK"] = np.einsum("btd,btk->dk", X, dK)
    g["WV"] = np.einsum("btd,btk->dk", X, dV)
    return g


def gradcheck():
    r = np.random.default_rng(0)
    P = init(r); X, y, *_ = make_batch(5, r)
    _, c = forward(P, X, y); g = backward(P, X, y, c)
    worst = 0
    for k in P:
        for _ in range(5):
            idx = tuple(r.integers(0, n) for n in P[k].shape)
            old = P[k][idx]; h = 1e-5
            P[k][idx] = old + h; lp, _ = forward(P, X, y)
            P[k][idx] = old - h; lm, _ = forward(P, X, y)
            P[k][idx] = old
            num = (lp - lm) / (2 * h)
            worst = max(worst, abs(num - g[k][idx]) / max(1e-8, abs(num) + abs(g[k][idx])))
    return worst


def attention_rows(P, sample):
    X, y = sample[0], sample[1]
    _, c = forward(P, X, y)
    return c["a"], c["p"]


def main():
    print("gradcheck worst rel err:", gradcheck())
    P = init(rng)
    m = {k: np.zeros_like(v) for k, v in P.items()}; v2 = {k: np.zeros_like(v) for k, v in P.items()}
    lr, b1, b2 = 3e-3, 0.9, 0.999
    test = make_batch(2000, np.random.default_rng(123))
    show = make_batch(3, np.random.default_rng(2026))
    snapshots, curve = {}, []
    steps = 3000
    for t in range(steps + 1):
        if t in (0, 300, steps):
            a, p = attention_rows(P, show)
            snapshots[t] = {"attn": a.round(3).tolist(),
                            "pred": p.argmax(1).tolist()}
        if t % 100 == 0:
            loss, c = forward(P, test[0], test[1])
            acc = (c["p"].argmax(1) == test[1]).mean()
            curve.append((t, round(float(loss), 3), round(float(acc), 3)))
        if t == steps:
            break
        X, y, *_ = make_batch(128, rng)
        _, c = forward(P, X, y); g = backward(P, X, y, c)
        for k in P:
            m[k] = b1 * m[k] + (1 - b1) * g[k]; v2[k] = b2 * v2[k] + (1 - b2) * g[k] ** 2
            mh = m[k] / (1 - b1 ** (t + 1)); vh = v2[k] / (1 - b2 ** (t + 1))
            P[k] -= lr * mh / (np.sqrt(vh) + 1e-8)

    # Which part of the token does each projection listen to?
    def sens(W, E):  # mean norm of projected (centered) embeddings
        Ec = E - E.mean(0)
        return float(np.linalg.norm(Ec @ W, axis=1).mean())
    ratios = {}
    for k in ("WQ", "WK", "WV"):
        ratios[k] = {"letter": sens(P[k], E_let), "digit": sens(P[k], E_dig[:QMARK])}
    _, lets, digs, qlet = show[1], show[2], show[3], show[4]
    out = {
        "curve": curve,
        "snapshots": snapshots,
        "show": {"letters": [[LETTERS[i] for i in row] for row in lets],
                 "digits": [[DIGITS[i] for i in row] for row in digs],
                 "query": [LETTERS[i] for i in qlet],
                 "answer": [DIGITS[i] for i in show[1]]},
        "sensitivity": ratios,
    }
    json.dump(out, open("toy_attention_result.json", "w"), indent=1)
    for row in curve[::3] + [curve[-1]]:
        print(row)
    print(json.dumps(out["sensitivity"], indent=1))
    for i in range(3):
        toks = [a + b for a, b in zip(out["show"]["letters"][i], out["show"]["digits"][i])] + [out["show"]["query"][i] + "?"]
        print(toks, "answer", out["show"]["answer"][i])
        for t in (0, 300, steps):
            print("  step", t, snapshots[t]["attn"][i], "pred", snapshots[t]["pred"][i])


if __name__ == "__main__":
    main()
