#!/usr/bin/env python3
"""A complete, deliberately tiny causal decoder using only Python's standard library.

Vocabulary: The=0, cat=1, sat=2, .=3. The tokenizer is whitespace splitting,
so punctuation must be separated by a space. This is NOT a real tokenizer.
There is one block, one 2-coordinate head, model width 2, and FFN width 3.
All coefficients are invented and untrained. A readable continuation is not
evidence that these weights learned language or that a real LLM works this way
semantically. The arithmetic operations and causal cache reuse are the lesson.

Run: python3 01-llm-internals-tiny-decoder.py
"""

import math
from dataclasses import dataclass, field


VOCAB = ("The", "cat", "sat", ".")
TOKEN_IDS = {token: index for index, token in enumerate(VOCAB)}
WIDTH = 2
FFN_WIDTH = 3
EPSILON = 1e-6

# Rows select token IDs. Positions have their own invented absolute table.
EMBEDDINGS = [[1.0, 0.0], [0.0, 1.0], [1.0, 0.2], [-1.0, 0.0]]
POSITIONS = [[0.0, 0.0], [0.1, -0.1], [0.2, -0.2], [0.3, -0.3]]

# Row-vector convention: [1 x input_width] @ [input_width x output_width].
W_Q = [[0.5, 0.0], [0.0, 0.5]]
W_K = [[1.0, 0.2], [-0.3, 1.0]]
W_V = [[1.0, 0.0], [0.0, 0.5]]
W_O = [[0.5, 0.1], [0.2, 0.5]]
W_UP = [[1.0, -1.0, 0.5], [0.5, 1.0, -0.5]]
W_DOWN = [[0.2, 0.0], [0.0, 0.2], [0.1, -0.1]]
W_LM = [[-0.5, 0.2, 0.0, 1.0], [-0.5, 0.1, 1.0, 0.0]]
ATTENTION_SCALE = [1.0, 1.0]
FFN_SCALE = [1.0, 1.0]
FINAL_SCALE = [1.0, 1.0]


def tokenize(text):
    """Map this teaching vocabulary's space-separated pieces to IDs."""
    return [TOKEN_IDS[piece] for piece in text.split()]


def decode(ids):
    """Keep spaces, including before punctuation, so all token boundaries show."""
    return " ".join(VOCAB[token_id] for token_id in ids)


def project(vector, matrix):
    assert vector and len(vector) == len(matrix)
    assert matrix[0] and all(len(row) == len(matrix[0]) for row in matrix)
    return [sum(value * matrix[i][j] for i, value in enumerate(vector))
            for j in range(len(matrix[0]))]


def add(left, right):
    assert len(left) == len(right)
    return [a + b for a, b in zip(left, right)]


def rmsnorm(vector, scale):
    assert vector and len(vector) == len(scale)
    divisor = math.sqrt(sum(value * value for value in vector) / len(vector)
                        + EPSILON)
    return [gain * value / divisor for gain, value in zip(scale, vector)]


def softmax(scores):
    maximum = max(scores)
    exponentials = [math.exp(score - maximum) for score in scores]
    total = sum(exponentials)
    return [value / total for value in exponentials]


def input_at(token_id, position):
    if not 0 <= token_id < len(VOCAB):
        raise ValueError("Token ID is outside the toy vocabulary")
    if not 0 <= position < len(POSITIONS):
        raise ValueError("The toy model has only four position vectors")
    return add(EMBEDDINGS[token_id], POSITIONS[position])


def attention(query, allowed_keys, allowed_values):
    assert allowed_keys and len(allowed_keys) == len(allowed_values)
    scores = [sum(q * k for q, k in zip(query, key)) / math.sqrt(WIDTH)
              for key in allowed_keys]
    weights = softmax(scores)
    mixed = [sum(weight * value[j]
                 for weight, value in zip(weights, allowed_values))
             for j in range(WIDTH)]
    return scores, weights, mixed


def finish_position(x, mixed):
    """W_O -> residual -> pre-norm FFN -> residual -> final norm -> LM head."""
    attention_update = project(mixed, W_O)
    residual = add(x, attention_update)
    normalized = rmsnorm(residual, FFN_SCALE)
    expanded = project(normalized, W_UP)
    activated = [max(0.0, value) for value in expanded]  # ReLU, coordinate-wise
    ffn_update = project(activated, W_DOWN)
    block_output = add(residual, ffn_update)
    final = rmsnorm(block_output, FINAL_SCALE)
    logits = project(final, W_LM)
    probabilities = softmax(logits)
    return {
        "attention_update": attention_update,
        "attention_residual": residual,
        "ffn_normalized": normalized,
        "ffn_expanded": expanded,
        "ffn_activated": activated,
        "ffn_update": ffn_update,
        "block_output": block_output,
        "final": final,
        "logits": logits,
        "probabilities": probabilities,
        "selected_id": max(range(len(VOCAB)), key=logits.__getitem__),
    }


def full_forward(ids):
    """Process all known positions, with each attention row sliced causally."""
    if not ids:
        raise ValueError("Supply at least one prompt token")
    inputs = [input_at(token_id, position) for position, token_id in enumerate(ids)]
    normalized = [rmsnorm(x, ATTENTION_SCALE) for x in inputs]
    queries = [project(x, W_Q) for x in normalized]
    keys = [project(x, W_K) for x in normalized]
    values = [project(x, W_V) for x in normalized]
    traces = []
    for position, (x, query) in enumerate(zip(inputs, queries)):
        # The mask includes the current position and excludes every future one.
        scores, weights, mixed = attention(query, keys[:position + 1],
                                           values[:position + 1])
        trace = finish_position(x, mixed)
        trace.update(x=x, normalized=normalized[position], query=query,
                     key=keys[position], value=values[position], scores=scores,
                     attention_weights=weights, mixed=mixed)
        traces.append(trace)
    return traces


@dataclass
class KVCache:
    # Exactly one block here. A deeper model needs a separate cache per block.
    keys: list = field(default_factory=list)
    values: list = field(default_factory=list)


def prefill(ids):
    """Process the prompt and retain this block's projected K/V."""
    traces = full_forward(ids)
    cache = KVCache([trace["key"][:] for trace in traces],
                    [trace["value"][:] for trace in traces])
    return traces, cache


def decode_step(token_id, cache):
    """Process ONE selected token at its next position, then predict its successor.

    Only append tokens from the same unchanged model/prefix to this cache.
    The token selected by this call is not itself in the cache yet.
    """
    assert len(cache.keys) == len(cache.values)
    position = len(cache.keys)
    x = input_at(token_id, position)
    normalized = rmsnorm(x, ATTENTION_SCALE)
    query = project(normalized, W_Q)
    key, value = project(normalized, W_K), project(normalized, W_V)
    cache.keys.append(key)
    cache.values.append(value)
    scores, weights, mixed = attention(query, cache.keys, cache.values)
    trace = finish_position(x, mixed)
    trace.update(x=x, normalized=normalized, query=query, key=key, value=value,
                 scores=scores, attention_weights=weights, mixed=mixed)
    return trace


def assert_vectors_close(left, right):
    assert len(left) == len(right)
    assert all(math.isclose(a, b, rel_tol=1e-12, abs_tol=1e-12)
               for a, b in zip(left, right))


def head_only_training_step(hidden, target_id, learning_rate=0.1):
    """One real gradient update of a COPY of W_LM, keeping this hidden vector fixed.

    This is not full-model backpropagation. The embeddings, attention, norms and
    FFN stay frozen. Return the new head rather than mutating our example model.
    """
    before = softmax(project(hidden, W_LM))
    logit_gradients = [p - float(i == target_id) for i, p in enumerate(before)]
    gradients = [[value * gradient for gradient in logit_gradients]
                 for value in hidden]
    updated = [[weight - learning_rate * gradient
                for weight, gradient in zip(row, gradient_row)]
               for row, gradient_row in zip(W_LM, gradients)]
    after = softmax(project(hidden, updated))
    return updated, gradients, -math.log(before[target_id]), -math.log(after[target_id])


def main():
    ids = tokenize("The cat")
    traces, cache = prefill(ids)
    last = traces[-1]
    print("Toy vocabulary:", ", ".join(f"{word}={i}" for i, word in enumerate(VOCAB)))
    print("Prompt IDs:", ids)
    print("Trace at 'cat' (position 1; positions count from zero):")
    labels = [
        ("embedding + position", "x"), ("attention RMSNorm", "normalized"),
        ("query", "query"), ("key", "key"), ("value", "value"),
        ("scaled scores", "scores"), ("attention weights", "attention_weights"),
        ("weighted values", "mixed"), ("W_O update", "attention_update"),
        ("first residual", "attention_residual"), ("FFN RMSNorm", "ffn_normalized"),
        ("FFN expansion", "ffn_expanded"), ("ReLU", "ffn_activated"),
        ("FFN update", "ffn_update"), ("second residual", "block_output"),
        ("final RMSNorm", "final"), ("vocabulary logits", "logits"),
        ("next-token probabilities", "probabilities"),
    ]
    for label, key in labels:
        print(f"  {label}: [" + ", ".join(f"{value:.6f}" for value in last[key]) + "]")
    selected = last["selected_id"]
    print(f"Greedy token after prefill: {VOCAB[selected]} (ID {selected})")
    cached = decode_step(selected, cache)
    recomputed = full_forward(ids + [selected])[-1]
    for key in ("block_output", "final", "logits", "probabilities"):
        assert_vectors_close(cached[key], recomputed[key])
    assert cached["selected_id"] == recomputed["selected_id"]
    print("Full-prefix and cached next step agree (tolerance 1e-12).")
    print("Next-step probabilities: [" + ", ".join(f"{p:.6f}" for p in cached["probabilities"]) + "]")
    next_id = cached["selected_id"]
    print(f"Next greedy token: {VOCAB[next_id]} (ID {next_id})")
    print("Displayed tokens:", decode(ids + [selected, next_id]))
    print(f"Cache contains {len(cache.keys)} processed positions; the last selected token is not processed.")
    _, _, old_loss, new_loss = head_only_training_step(last["final"], TOKEN_IDS["sat"])
    print(f"Head-only training on target 'sat': loss {old_loss:.6f} -> {new_loss:.6f}.")
    print("Only a copy of W_LM was updated; this is not full-model backpropagation.")


if __name__ == "__main__":
    main()
