# /// script
# requires-python = ">=3.12"
# dependencies = ["marimo==0.25.1", "numpy==2.5.3"]
# ///
# Generated by scripts/learning/export-masked-readout.ts. Edit that source to regenerate.
import marimo

__generated_with = "0.25.1"
app = marimo.App(width="medium")

@app.cell
def _():
    import marimo as mo
    return (mo,)

@app.cell
def _(mo):
    mo.md("""
    # Which tokens teach this readout?
    Change one declaration and inspect one update from zero weights.
    The hidden rows are authored and frozen. This is a numerical teaching example,
    not a trained transformer or a measurement of generalization.
    [Return to the lesson](https://continuousfunction.ai/ai-lab/transformers/from-text-to-update/#masked-readout-update).
    """)
    return

@app.cell
def _(mo):
    policy = mo.ui.dropdown({"Completion targets": "completion", "All valid next-token targets": "all"}, value="Completion targets", label="Targets that count")
    grouping = mo.ui.dropdown({"One batch": "one-batch", "Split, divide once": "split-token-mean", "Split, average means": "split-mean-of-means"}, value="One batch", label="Grouping")
    mo.vstack([policy, grouping])
    return (policy, grouping)

@app.cell
def _():
    def experiment(policy, grouping):
        if policy not in ("completion", "all") or grouping not in ("one-batch", "split-token-mean", "split-mean-of-means"):
            raise ValueError("Choose a declared target policy and grouping")
        import numpy as np

        # Output classes 0 red, 1 blue, 2 green (V = 3). PAD marks padding; it is not a class.
        RED, BLUE, GREEN, PAD = 0, 1, 2, "PAD"
        V, d, eta = 3, 2, 0.3

        # Two right-padded sequences, B = 2 and T = 4. Roles mark the completion; they are not tokens.
        tokens = [[GREEN, GREEN, RED, RED],
                  [GREEN, BLUE, PAD, PAD]]
        roles = [["prompt", "prompt", "completion", "completion"],
                 ["prompt", "completion", None, None]]

        # Authored, frozen hidden rows, one per input position: shape (B, T, d).
        H = np.array([[[0.0, 1.0], [1.0, 0.0], [1.0, 0.0], [0.0, 0.0]],
                      [[1.0, 0.0], [0.0, 0.0], [0.0, 0.0], [0.0, 0.0]]])
        W0, b0 = np.zeros((V, d)), np.zeros(V)


        def align(policy):
            """Shift first: row (i, t) is scored on token t + 1 of the same sequence.
            Returns targets y (a class, or None) and eligibility m (0 or 1), each B x T."""
            y, m = [], []
            for seq, role in zip(tokens, roles):
                nxt, nxt_role = seq[1:] + [None], role[1:] + [None]  # the last row has no successor
                y.append([None if tok in (None, PAD) else tok for tok in nxt])
                m.append([int(tok not in (None, PAD) and (policy == "all" or r == "completion"))
                          for tok, r in zip(nxt, nxt_role)])
            return y, m


        def finite(*values):
            """Reject any result float64 cannot hold, instead of carrying inf or nan onward."""
            for value in values:
                if not np.isfinite(value).all():
                    raise FloatingPointError("a result is not representable in float64")
            return values


        def sums(seqs, y, m, W, b):
            """Summed loss, eligible count and summed gradients over the rows of the given sequences.
            Only rows with m = 1 are read; every other row has delta = 0."""
            S, M, GW, Gb = 0.0, 0, np.zeros(W.shape), np.zeros(b.shape)
            for i in seqs:
                rows = [t for t in range(len(y[i])) if m[i][t] == 1]
                if not rows:
                    continue
                Hr = np.array([H[i, t].tolist() for t in rows])   # (R, d) eligible hidden rows
                Z = Hr @ W.T + b                                   # (R, d) @ (d, V) + (V,) -> (R, V) logits
                Zs = Z - Z.max(axis=1, keepdims=True)              # shifted: each row's largest is 0
                finite(Z, Zs)
                E = np.exp(Zs)
                total = E.sum(axis=1, keepdims=True)               # (R, 1), between 1 and V
                P = E / total                                      # (R, V) probabilities
                Y = np.array([[float(y[i][t] == c) for c in range(len(b))] for t in rows])  # one-hot targets
                # Loss in shifted coordinates: log sum exp(z - max) - (z_y - max). The offset never returns.
                S += float((np.log(total) - (Zs * Y).sum(axis=1, keepdims=True)).sum())
                M += len(rows)
                R = P - Y                                          # (R, V): p - e_y
                GW = GW + R.T @ Hr                                 # (V, R) @ (R, d) = sum of (p - e_y) h^T
                Gb = Gb + R.sum(axis=0)
                finite(S, GW, Gb)                                  # summed, not yet divided
            return S, M, GW, Gb


        def token_mean(parts, y, m, W, b):
            """The declared objective: add every microbatch's sums and counts, then divide once."""
            S, M, GW, Gb = 0.0, 0, np.zeros(W.shape), np.zeros(b.shape)
            for part in parts:
                s, k, gw, gb = sums(part, y, m, W, b)
                S, M, GW, Gb = S + s, M + k, GW + gw, Gb + gb
            finite(S, GW, Gb)
            if M == 0:
                return None                                        # no supervised target: no mean, no step
            return S / M, GW / M, Gb / M


        def mean_of_means(parts, y, m, W, b):
            """A different objective: each microbatch's own mean, averaged with equal weight."""
            pieces = [sums(part, y, m, W, b) for part in parts]
            if any(k == 0 for _, k, _, _ in pieces):
                return None                                        # an empty microbatch has no mean
            finite(sum(s for s, _, _, _ in pieces))                # the summed raw loss, as the ledger reports it
            K = len(pieces)                                        # each target weighs 1 / (K k)
            return finite(sum(s / k / K for s, k, _, _ in pieces),
                          sum(gw / k / K for _, k, gw, _ in pieces),
                          sum(gb / k / K for _, k, _, gb in pieces))


        def sgd_step(W, b, gW, gb, eta):
            """One plain step. A parameter float64 cannot hold is rejected, never clamped."""
            return finite(W - eta * gW, b - eta * gb)


        def readout(W, b, h):
            """Shifted logits, probabilities and log sum exp(z - max) for one hidden vector."""
            z = W @ h + b                                          # (V,)
            zs = z - z.max()
            finite(z, zs)
            e = np.exp(zs)
            return zs, e / e.sum(), float(np.log(e.sum()))


        def loss_at(W, b, h, target):
            zs, _, log_total = readout(W, b, h)
            return log_total - float(zs[target])


        def show(a):
            return (np.round(a, 6) + 0.0).tolist()                 # + 0.0 prints -0.0 as 0.0


        REFERENCE_CASE = {"example": "masked-readout-one-step-v1", "policy": "completion",
                          "grouping": "one-batch", "eta": 0.3}


        def reproduce_case(case):
            """This version fixes tokens, H, zero W/b and one step. Export only its declared controls."""
            if (not isinstance(case, dict) or set(case) != set(REFERENCE_CASE)
                    or case["example"] != REFERENCE_CASE["example"]
                    or case["policy"] not in ("completion", "all")
                    or case["grouping"] not in ("one-batch", "split-token-mean", "split-mean-of-means")
                    or type(case["eta"]) not in (int, float) or case["eta"] != eta):
                raise ValueError("This masked-readout case is not supported by this code")
            y, m = align(case["policy"])
            parts = [[0, 1]] if case["grouping"] == "one-batch" else [[0], [1]]
            reduction = mean_of_means if case["grouping"] == "split-mean-of-means" else token_mean
            L, gW, gb = reduction(parts, y, m, W0, b0)
            W1, b1 = sgd_step(W0, b0, gW, gb, case["eta"])
            before, _, _ = token_mean([[0, 1]], y, m, W0, b0)
            after, _, _ = token_mean([[0, 1]], y, m, W1, b1)
            h = np.array([1.0, 0.0])
            _, p1, _ = readout(W1, b1, h)
            return {**case, "y": y, "m": m, "targets": sum(map(sum, m)),
                    "objective": float(L), "numerator": float(sums([0, 1], y, m, W0, b0)[0]),
                    "training_before": float(before), "training_after": float(after),
                    "grad_W": gW.tolist(), "grad_b": gb.tolist(), "W_after": W1.tolist(), "b_after": b1.tolist(),
                    "row_loss_after": [float(loss_at(W1, b1, H[i, t], y[i][t]))
                                       for part in parts for i in part for t in range(len(y[i])) if m[i][t]],
                    "held_out_before": float(loss_at(W0, b0, h, BLUE)),
                    "held_out_after": float(loss_at(W1, b1, h, BLUE)), "probabilities_after": p1.tolist()}


        return reproduce_case({**REFERENCE_CASE, "policy": policy, "grouping": grouping})
    return (experiment,)

@app.cell
def _(mo, experiment, policy, grouping, result_json):
    result = experiment(policy.value, grouping.value)
    mo.vstack([
        mo.md(f"**{result['targets']} eligible targets.** Mean loss on those targets: {result['training_before']:.6f} → {result['training_after']:.6f}."),
        mo.ui.table([{"class": label, "gradient b": f"{result['grad_b'][i]:.6f}".replace("-0.000000", "0.000000"), "updated b": f"{result['b_after'][i]:.6f}".replace("-0.000000", "0.000000")} for i, label in enumerate(["red", "blue", "green"])], selection=None),
        mo.md("Table values are rounded to six decimals; the inspected and downloaded values keep full precision."),
        mo.md(f"Blue training row and authored blue counterexample after the update: **{result['held_out_after']:.6f}**. Both have the same hidden row (1, 0) and blue target; the readout cannot distinguish their different prompts. This loss may worsen while the training mean falls."),
        mo.accordion({"Inspect every returned value": mo.json(result)}),
        mo.download(result_json(result), filename="readout-result.json", label="Download this calculation")
    ])
    return (result,)

@app.cell
def _():
    def result_json(result):
        import json
        return json.dumps(result, indent=2, allow_nan=False).encode()
    return (result_json,)

if __name__ == "__main__":
    app.run()
