from __future__ import annotations

import importlib.util
import json
import math
import sys
from pathlib import Path

import torch
from torch.nn import functional as F


TRAINING_SCRIPT = Path(__file__).with_name("train-tiny-transformer.py")
spec = importlib.util.spec_from_file_location("tiny_training", TRAINING_SCRIPT)
training = importlib.util.module_from_spec(spec)
assert spec.loader is not None
sys.modules[spec.name] = training
spec.loader.exec_module(training)


TOKENS = ["<bos>", "A", "B", "C", "D", "X", "Y", "<eos>"]
EOS_ID = 7


def split_heads(records: torch.Tensor, heads: int) -> torch.Tensor:
    batch, length, width = records.shape
    return records.reshape(batch, length, heads, width // heads).transpose(1, 2)


def join_heads(records: torch.Tensor) -> torch.Tensor:
    batch, heads, length, head_width = records.shape
    return records.transpose(1, 2).reshape(batch, length, heads * head_width)


@torch.inference_mode()
def cached_token(
    model,
    token_id: int,
    position: int,
    caches: list[tuple[torch.Tensor, torch.Tensor] | None],
    wrong_position: bool = False,
    wrong_layer_cache: bool = False,
) -> tuple[torch.Tensor, list[tuple[torch.Tensor, torch.Tensor]]]:
    if position >= model.config.context_length:
        raise ValueError("cache would exceed the context length")
    used_position = 0 if wrong_position else position
    ids = torch.tensor([[token_id]], dtype=torch.long)
    hidden = model.token_embedding(ids) + model.position_embedding[used_position]
    updated = []

    for layer_index, block in enumerate(model.blocks):
        normalized = block.norm1(hidden)
        qkv = F.linear(normalized, block.attention.in_proj_weight)
        query, key, value = qkv.chunk(3, dim=-1)
        query = split_heads(query, model.config.n_heads)
        key = split_heads(key, model.config.n_heads)
        value = split_heads(value, model.config.n_heads)

        previous = caches[0] if wrong_layer_cache else caches[layer_index]
        if previous is not None:
            key = torch.cat((previous[0], key), dim=2)
            value = torch.cat((previous[1], value), dim=2)
        updated.append((key, value))

        scores = query @ key.transpose(-2, -1) / math.sqrt(
            model.config.d_model // model.config.n_heads
        )
        weights = scores.softmax(dim=-1)
        attended = block.attention.out_proj(join_heads(weights @ value))
        hidden = hidden + attended
        hidden = hidden + block.mlp(block.norm2(hidden))

    logits = model.readout(model.final_norm(hidden))[:, -1]
    return logits, updated


@torch.inference_mode()
def cached_prefill(
    model,
    prompt: list[int],
    wrong_position: bool = False,
    wrong_layer_cache: bool = False,
) -> tuple[torch.Tensor, list[tuple[torch.Tensor, torch.Tensor]]]:
    caches = [None] * model.config.n_layers
    logits = None
    for position, token_id in enumerate(prompt):
        logits, caches = cached_token(
            model,
            token_id,
            position,
            caches,
            wrong_position=wrong_position,
            wrong_layer_cache=wrong_layer_cache,
        )
    assert logits is not None
    return logits, caches


@torch.inference_mode()
def full_last_logits(model, prompt: list[int]) -> torch.Tensor:
    ids = torch.tensor([prompt], dtype=torch.long)
    return model(ids)[:, -1]


def choose_token(
    logits: torch.Tensor,
    policy: str,
    generator: torch.Generator | None = None,
    temperature: float = 1.0,
    top_k: int | None = None,
) -> int:
    row = logits[0]
    if policy == "greedy" or temperature == 0:
        return int(row.argmax())
    if temperature < 0:
        raise ValueError("temperature must be non-negative")
    scaled = row / temperature
    if top_k is not None:
        if not 1 <= top_k <= len(row):
            raise ValueError("top_k must lie between 1 and vocabulary size")
        threshold = torch.topk(scaled, top_k).values[-1]
        scaled = scaled.masked_fill(scaled < threshold, float("-inf"))
    probabilities = scaled.softmax(dim=-1)
    return int(torch.multinomial(probabilities, 1, generator=generator))


@torch.inference_mode()
def generate(
    model,
    prompt: list[int],
    policy: str,
    seed: int | None = None,
    temperature: float = 1.0,
    top_k: int | None = None,
) -> dict:
    generated = list(prompt)
    logits, caches = cached_prefill(model, prompt)
    generator = None if seed is None else torch.Generator().manual_seed(seed)
    stop = ""
    while True:
        token_id = choose_token(
            logits,
            policy,
            generator=generator,
            temperature=temperature,
            top_k=top_k,
        )
        generated.append(token_id)
        if token_id == EOS_ID:
            stop = "eos"
            break
        if caches[0][0].shape[2] >= model.config.context_length:
            stop = "context"
            break
        position = caches[0][0].shape[2]
        logits, caches = cached_token(model, token_id, position, caches)
    return {
        "ids": generated,
        "tokens": [TOKENS[token_id] for token_id in generated],
        "stop": stop,
        "cached_positions": int(caches[0][0].shape[2]),
    }


def main() -> None:
    config = training.Config()
    training_result, state = training.train(config, return_state=True)
    torch.manual_seed(config.model_seed + 2000)
    model = training.TinyDecoder(config)
    model.load_state_dict(state)
    model.eval()

    maximum_error = 0.0
    prefix_checks = []
    for sequence in training.SEQUENCES:
        input_ids = list(sequence[:-1])
        for length in range(1, len(input_ids) + 1):
            prompt = input_ids[:length]
            full = full_last_logits(model, prompt)
            cached, caches = cached_prefill(model, prompt)
            error = float((full - cached).abs().max())
            maximum_error = max(maximum_error, error)
            expected_shape = (1, 2, length, 2)
            assert all(tuple(key.shape) == expected_shape for key, _ in caches)
            assert all(tuple(value.shape) == expected_shape for _, value in caches)
            prefix_checks.append(
                {
                    "prompt": [TOKENS[token_id] for token_id in prompt],
                    "maximum_logit_error": error,
                    "cache_shape_per_layer": expected_shape,
                }
            )
    tolerance = 1e-5
    assert maximum_error < tolerance

    correct, _ = cached_prefill(model, [0, 1])
    wrong, _ = cached_prefill(model, [0, 1], wrong_position=True)
    wrong_position_error = float((correct - wrong).abs().max())
    assert wrong_position_error > tolerance
    wrong_layer, _ = cached_prefill(model, [0, 1], wrong_layer_cache=True)
    wrong_layer_error = float((correct - wrong_layer).abs().max())
    assert wrong_layer_error > tolerance

    greedy = generate(model, [0], "greedy")
    sampled = generate(
        model,
        [0],
        "sample",
        seed=41,
        temperature=1.0,
        top_k=3,
    )

    length = config.context_length
    cache_scalars = 2 * config.n_layers * 1 * config.n_heads * length * (
        config.d_model // config.n_heads
    )
    result = {
        "torch_version": torch.__version__,
        "training_final_validation_loss": training_result["final_validation_loss"],
        "prefix_checks": prefix_checks,
        "maximum_logit_error": maximum_error,
        "tolerance": tolerance,
        "wrong_position_error": wrong_position_error,
        "wrong_layer_cache_error": wrong_layer_error,
        "greedy": greedy,
        "sampled_seed_41_temperature_1_top_k_3": sampled,
        "cache_at_context_limit": {
            "shape_per_tensor_per_layer": [1, 2, 4, 2],
            "scalars": cache_scalars,
            "float32_bytes": cache_scalars * 4,
        },
    }
    print(json.dumps(result, indent=2))


if __name__ == "__main__":
    main()
