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()}


# Standalone reference demonstration; a supplied case is reproduced below instead.
if "case" not in globals():
    y, m = align("completion")
    L, gW, gb = token_mean([[0, 1]], y, m, W0, b0)
    W1, b1 = sgd_step(W0, b0, gW, gb, eta)                     # one SGD step
    L1, _, _ = token_mean([[0, 1]], y, m, W1, b1)
    h = np.array([1.0, 0.0])
    _, p1, _ = readout(W1, b1, h)

    print("targets M =", sum(map(sum, m)), "| loss", round(L, 6), "->", round(L1, 6))
    print("grad_W", show(gW), "grad_b", show(gb))
    print("W'", show(W1), "b'", show(b1))
    print("p' at h = (1, 0):", show(p1))

    # The same rows in two microbatches, then one step: add sums, divide once.
    _, gW_split, gb_split = token_mean([[0], [1]], y, m, W0, b0)
    assert np.allclose(gW_split, gW) and np.allclose(gb_split, gb)
    # Averaging the two microbatch means starts at the same loss but moves differently.
    L_mm, gW_mm, gb_mm = mean_of_means([[0], [1]], y, m, W0, b0)
    print("mean of means: loss", round(L_mm, 6), "grad_b", show(gb_mm))

    # A held-out case declared in advance: red prompt -> blue completion, also h = (1, 0).
    print("held-out loss", round(loss_at(W0, b0, h, BLUE), 6), "->", round(loss_at(W1, b1, h, BLUE), 6))

    # The other declared target set, and a batch with nothing to learn from.
    y_all, m_all = align("all")
    L_all, gW_all, gb_all = token_mean([[0, 1]], y_all, m_all, W0, b0)
    print("all targets: M =", sum(map(sum, m_all)), "grad_b", show(gb_all))
    assert token_mean([[0, 1]], [[None] * 4] * 2, [[0] * 4] * 2, W0, b0) is None

# The browser sets case from JSON. Locally, use runpy.run_path(..., init_globals={"case": data}).
# Full-precision JSON is the comparison contract; the preceding reference printout is rounded.
import json
result = reproduce_case(globals().get("case", REFERENCE_CASE))
if "case" in globals():
    print("Case:", result["policy"], "|", result["grouping"], "| one step, eta =", result["eta"])
    print("targets M =", result["targets"], "| training mean", result["training_before"], "->", result["training_after"])
    print("grad_W", result["grad_W"], "grad_b", result["grad_b"])
    print("W'", result["W_after"], "b'", result["b_after"])
    print("p' at h = (1, 0):", result["probabilities_after"])
print("MASKED_READOUT_CASE=" + json.dumps(result, allow_nan=False, separators=(",", ":")))
