"""
Lava Leap: a 40-tile toy game for learning GRPO from scratch.

A delivery bot crosses a foundry floor (tiles 0..39). Molten-lava channels
block some tiles. At every step the bot sees the next three tiles and picks
STEP (exactly 1 tile) or LEAP (a coin flip: 2 or 3 tiles). The only feedback
it ever gets is one number at the very end: how far it got.

Everything here is NumPy: the game, an eight-weight policy, three ways to
train it (REINFORCE, a toy PPO with a four-weight critic, and GRPO), and an
exact evaluator. The game, policy and update sections (between the
"# ---- [name]" dividers) are quoted verbatim in the blog post.
"""
from collections import namedtuple

import numpy as np

# ---- [game] -----------------------------------------------------------
FLOOR, EXIT = 40, 39        # tiles 0..39; the exit dock is tile 39
STEP, LEAP = 0, 1           # the bot's two moves

def new_floor(rng):
    """3-7 lava channels, 1-2 tiles wide, in tiles 5..36, >= 2 apart."""
    k = rng.integers(3, 8)                    # how many channels
    widths = rng.integers(1, 3, size=k)       # each 1 or 2 tiles wide
    spare = 32 - widths.sum() - 2 * (k - 1)   # leftover safe tiles
    gaps = rng.multinomial(spare, np.ones(k + 1) / (k + 1))
    lava, t = np.zeros(FLOOR, dtype=bool), 5 + gaps[0]
    for w, extra in zip(widths, gaps[1:]):
        lava[t:t + w] = True
        t += w + 2 + extra                    # 2 safe tiles + any spare
    return lava

def sensors(lava, pos):
    """What the bot sees: lava at pos+1, pos+2, pos+3, then a bias 1.0."""
    ahead = [float(t < FLOOR and lava[t]) for t in range(pos + 1, pos + 4)]
    return np.array(ahead + [1.0])            # past tile 39 reads safe

def step(pos, action, rng):
    """STEP moves exactly 1 tile. LEAP is a coin flip: 2 or 3 tiles."""
    return pos + 1 if action == STEP else pos + int(rng.integers(2, 4))

def score(pos):
    """The only feedback, given once at the end: how far the bot got."""
    return min(pos, EXIT) / EXIT

# ---- [policy] ---------------------------------------------------------
Run = namedtuple("Run", "states actions path reward")

def policy(W, s):
    """pi(.|s) = softmax(s @ W). W is 4x2: one column of weights per move.
    Accepts one sensor row (shape 4) or a stack of rows (shape N x 4)."""
    z = s @ W
    e = np.exp(z - z.max(axis=-1, keepdims=True))   # max-shift for safety
    return e / e.sum(axis=-1, keepdims=True)

def play(W, lava, rng):
    """One run: look, sample a move from pi, move; until lava or the exit."""
    pos, states, actions, path = 0, [], [], [0]
    while pos < EXIT and not lava[pos]:
        s = sensors(lava, pos)
        a = LEAP if rng.random() < policy(W, s)[LEAP] else STEP
        states.append(s)
        actions.append(a)
        pos = step(pos, a, rng)
        path.append(min(pos, EXIT))
    return Run(np.array(states), np.array(actions), path, score(pos))

# ---- [batch] ----------------------------------------------------------
def batch(runs, per_run):
    """Stack every step of every run; each step inherits its run's number."""
    S = np.concatenate([run.states for run in runs])      # N x 4 sensor rows
    acts = np.concatenate([run.actions for run in runs])  # the N moves taken
    vals = np.concatenate([np.full(len(run.actions), x)
                           for run, x in zip(runs, per_run)])
    return S, acts, vals

# ---- [reinforce] ------------------------------------------------------
def reinforce_update(W, runs, lr=0.5):
    """REINFORCE: push every move up in proportion to its run's raw reward."""
    S, acts, R = batch(runs, [run.reward for run in runs])
    push = np.eye(2)[acts] - policy(W, S)   # onehot(a) - pi = grad log pi
    return W + lr * S.T @ (R[:, None] * push) / len(acts)

# ---- [ppo] ------------------------------------------------------------
def ppo_update(W, v, runs, lr=0.5, eps=0.2, passes=4, lr_critic=0.5,
               trace=None):
    """Toy PPO: advantage = actual score - critic's forecast V(s) = v.s"""
    S, acts, R = batch(runs, [run.reward for run in runs])
    A = R - S @ v                           # did the run beat the forecast?
    W = clipped_update(W, S, acts, A, lr, eps, passes, trace)
    for _ in range(passes):                 # critic: fit forecasts to scores
        v = v + lr_critic * S.T @ (R - S @ v) / len(R)
    return W, v

# ---- [grpo] -----------------------------------------------------------
def grpo_update(W, runs, lr=0.5, eps=0.2, passes=4, trace=None):
    """GRPO: one floor played G times; each run graded against its group."""
    r = np.array([run.reward for run in runs])
    adv = (r - r.mean()) / (r.std() + 1e-8)   # grade on a curve: no critic
    S, acts, A = batch(runs, adv)             # every step shares its run's A
    return clipped_update(W, S, acts, A, lr, eps, passes, trace)

def clipped_update(W, S, acts, A, lr=0.5, eps=0.2, passes=4, trace=None):
    """PPO's machinery, unchanged in GRPO: reuse one batch for several
    passes, weight each step by rho = pi_new / pi_old, stop at the clip."""
    N, rows = len(acts), np.arange(len(acts))
    onehot = np.eye(2)[acts]
    pi_old = policy(W, S)[rows, acts]         # frozen: the policy that played
    for _ in range(passes):
        pi = policy(W, S)
        rho = pi[rows, acts] / pi_old         # exactly 1 on pass 1
        clipped = ((A > 0) & (rho > 1 + eps)) | ((A < 0) & (rho < 1 - eps))
        coef = np.where(clipped, 0.0, A * rho) / N   # clipped: no gradient
        # gradient = sum over steps of coef * outer(s, onehot(a) - pi)
        grad = S.T @ (coef[:, None] * (onehot - pi))
        if trace is not None:                 # the walkthrough reads these
            trace.append(dict(W=W.copy(), pi=pi, rho=rho, clipped=clipped,
                              coef=coef, grad=grad))
        W = W + lr * grad
    return W

# ---- [training] -----------------------------------------------------
def train(algo, seed, rounds=300, G=8, lr=0.5, eps=0.2, passes=4,
          eval_floors=None, eval_every=10, snapshots=(), clip_log=None):
    """Train one bot. Every round: a fresh floor, played G times, one update.
    All three algorithms see the same floor sequence for a given seed.
    clip_log (a list, optional) gets one dict per update: the batch size,
    the mean |advantage|, and how many steps were clipped on each pass."""
    floor_rng = np.random.default_rng([seed, 0])
    play_rng = np.random.default_rng([seed, 1])
    W, v = np.zeros((4, 2)), np.zeros(4)      # all zeros: every move a coin flip
    curve, saved = [], {}
    for t in range(rounds + 1):
        if eval_floors is not None and t % eval_every == 0:
            curve.append((t, *evaluate(W, eval_floors)))
        if t in snapshots:
            saved[t] = (W.copy(), v.copy())
        if t == rounds:
            break
        lava = new_floor(floor_rng)
        runs = [play(W, lava, play_rng) for _ in range(G)]
        trace = [] if clip_log is not None else None
        if algo == "reinforce":
            W = reinforce_update(W, runs, lr)
        elif algo == "ppo":
            W, v = ppo_update(W, v, runs, lr, eps, passes, trace=trace)
        elif algo == "grpo":
            W = grpo_update(W, runs, lr, eps, passes, trace=trace)
        else:
            raise ValueError(algo)
        if trace:                             # pass 1: rho = 1, so coef * N = A
            clip_log.append(dict(steps=len(trace[0]["rho"]),
                                 mean_abs_adv=float(np.abs(trace[0]["coef"]).sum()),
                                 clipped=[int(p["clipped"].sum()) for p in trace]))
    return W, v, curve, saved

# ---- [evaluation] ---------------------------------------------------
PATTERNS = np.array([[b >> 2 & 1, b >> 1 & 1, b & 1, 1] for b in range(8)],
                    dtype=float)

def pattern_index(s):
    """Sensor row(s) -> 0..7, reading [lava+1, lava+2, lava+3] as 3 bits."""
    s = np.asarray(s)
    return (4 * s[..., 0] + 2 * s[..., 1] + s[..., 2]).astype(int)

def leap_table(W):
    """P(LEAP) for each of the 8 sensor patterns: the policy as a lookup table."""
    return policy(W, PATTERNS)[:, LEAP]

def exact_outcomes(p_leap, floors):
    """Exact expected score and exit probability, per floor, for a policy given
    as an 8-entry P(LEAP) table. A backward pass over the tiles averages over
    every coin flip (the policy's and the leap's): no sampling noise."""
    lava = np.asarray(floors, dtype=bool)
    F = len(lava)
    ahead = np.concatenate([lava, np.zeros((F, 3), dtype=bool)], axis=1).astype(int)
    V = np.ones((F, FLOOR + 3))               # expected final score, per tile
    P = np.ones((F, FLOOR + 3))               # probability of reaching the exit
    for pos in range(EXIT - 1, -1, -1):
        pattern = 4 * ahead[:, pos + 1] + 2 * ahead[:, pos + 2] + ahead[:, pos + 3]
        pl = np.asarray(p_leap)[pattern]
        v_go = (1 - pl) * V[:, pos + 1] + pl * 0.5 * (V[:, pos + 2] + V[:, pos + 3])
        p_go = (1 - pl) * P[:, pos + 1] + pl * 0.5 * (P[:, pos + 2] + P[:, pos + 3])
        V[:, pos] = np.where(lava[:, pos], pos / EXIT, v_go)
        P[:, pos] = np.where(lava[:, pos], 0.0, p_go)
    return V[:, 0], P[:, 0]

def evaluate(W, floors):
    """Mean exact score and exit rate of the (sampled) policy W over floors."""
    v, p = exact_outcomes(leap_table(W), floors)
    return float(v.mean()), float(p.mean())

def all_deterministic_policies(floors):
    """Score all 2^8 = 256 deterministic policies (one fixed move per pattern).
    Returns (scores, exit_rates) by policy id; bit b of id = LEAP on pattern b."""
    scores, exits = np.zeros(256), np.zeros(256)
    for pid in range(256):
        table = np.array([(pid >> b) & 1 for b in range(8)], dtype=float)
        v, p = exact_outcomes(table, floors)
        scores[pid], exits[pid] = v.mean(), p.mean()
    return scores, exits

def simulate(p_leap, floors, plays, rng):
    """Monte Carlo twin of exact_outcomes: actually play each floor `plays` times."""
    scores, exits = [], []
    for lava in floors:
        for _ in range(plays):
            pos = 0
            while pos < EXIT and not lava[pos]:
                p = p_leap[pattern_index(sensors(lava, pos))]
                a = LEAP if rng.random() < p else STEP
                pos = step(pos, a, rng)
            scores.append(score(pos))
            exits.append(pos >= EXIT)
    return float(np.mean(scores)), float(np.mean(exits))
