from __future__ import annotations

import copy
import json
import math
import random
import tempfile
from dataclasses import asdict, dataclass
from pathlib import Path

import torch
from torch import nn
from torch.nn import functional as F


VOCAB = {"<bos>": 0, "A": 1, "B": 2, "C": 3, "D": 4, "X": 5, "Y": 6, "<eos>": 7}
SEQUENCES = (
    (0, 1, 2, 3, 7),
    (0, 1, 2, 4, 7),
    (0, 5, 6, 5, 7),
    (0, 6, 5, 6, 7),
)
PROBABILITIES = (0.25, 0.25, 0.25, 0.25)


@dataclass(frozen=True)
class Config:
    vocab_size: int = 8
    context_length: int = 4
    d_model: int = 4
    n_heads: int = 2
    n_layers: int = 2
    d_ff: int = 8
    init_std: float = 0.1
    model_seed: int = 7
    train_seed: int = 17
    validation_seed: int = 29
    train_size: int = 256
    validation_size: int = 1024
    batch_size: int = 32
    steps: int = 1000
    learning_rate: float = 0.01
    clip_norm: float | None = 1.0


class Block(nn.Module):
    def __init__(self, config: Config):
        super().__init__()
        self.norm1 = nn.LayerNorm(config.d_model)
        self.attention = nn.MultiheadAttention(
            config.d_model,
            config.n_heads,
            dropout=0.0,
            bias=False,
            batch_first=True,
        )
        self.norm2 = nn.LayerNorm(config.d_model)
        self.mlp = nn.Sequential(
            nn.Linear(config.d_model, config.d_ff, bias=True),
            nn.ReLU(),
            nn.Linear(config.d_ff, config.d_model, bias=True),
        )

    def forward(self, x: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:
        normalized = self.norm1(x)
        attended, _ = self.attention(
            normalized,
            normalized,
            normalized,
            attn_mask=mask,
            need_weights=False,
        )
        x = x + attended
        x = x + self.mlp(self.norm2(x))
        return x


class TinyDecoder(nn.Module):
    def __init__(self, config: Config):
        super().__init__()
        self.config = config
        self.token_embedding = nn.Embedding(config.vocab_size, config.d_model)
        self.position_embedding = nn.Parameter(
            torch.empty(config.context_length, config.d_model)
        )
        self.blocks = nn.ModuleList([Block(config) for _ in range(config.n_layers)])
        self.final_norm = nn.LayerNorm(config.d_model)
        self.readout = nn.Linear(config.d_model, config.vocab_size, bias=False)
        self.apply(self._initialize_module)
        nn.init.normal_(self.position_embedding, mean=0.0, std=config.init_std)
        self.readout.weight = self.token_embedding.weight

    def _initialize_module(self, module: nn.Module) -> None:
        if isinstance(module, (nn.Linear, nn.Embedding)):
            nn.init.normal_(module.weight, mean=0.0, std=self.config.init_std)
            if isinstance(module, nn.Linear) and module.bias is not None:
                nn.init.zeros_(module.bias)
        elif isinstance(module, nn.MultiheadAttention):
            nn.init.normal_(
                module.in_proj_weight,
                mean=0.0,
                std=self.config.init_std,
            )

    def forward(self, token_ids: torch.Tensor, causal: bool = True) -> torch.Tensor:
        _, length = token_ids.shape
        if length > self.config.context_length:
            raise ValueError("sequence exceeds the context length")
        hidden = self.token_embedding(token_ids) + self.position_embedding[:length]
        if causal:
            mask = torch.triu(
                torch.full((length, length), float("-inf"), device=token_ids.device),
                diagonal=1,
            )
        else:
            mask = torch.zeros((length, length), device=token_ids.device)
        for block in self.blocks:
            hidden = block(hidden, mask)
        return self.readout(self.final_norm(hidden))


def make_dataset(size: int, seed: int) -> tuple[torch.Tensor, torch.Tensor, list[int]]:
    generator = random.Random(seed)
    indices = generator.choices(range(len(SEQUENCES)), weights=PROBABILITIES, k=size)
    rows = torch.tensor([SEQUENCES[index] for index in indices], dtype=torch.long)
    counts = [indices.count(index) for index in range(len(SEQUENCES))]
    return rows[:, :-1], rows[:, 1:], counts


def token_loss(logits: torch.Tensor, targets: torch.Tensor) -> torch.Tensor:
    return F.cross_entropy(logits.reshape(-1, logits.shape[-1]), targets.reshape(-1))


@torch.no_grad()
def evaluate(model: TinyDecoder, inputs: torch.Tensor, targets: torch.Tensor) -> float:
    model.eval()
    return float(token_loss(model(inputs), targets))


def gradient_norm(model: nn.Module) -> float:
    total = torch.zeros(())
    for parameter in model.parameters():
        if parameter.grad is not None:
            total += parameter.grad.detach().pow(2).sum()
    return float(total.sqrt())


def train(
    config: Config,
    one_sequence: bool = False,
    return_state: bool = False,
) -> dict | tuple[dict, dict]:
    torch.manual_seed(config.model_seed)
    model = TinyDecoder(config)
    parameter_count = sum(parameter.numel() for parameter in model.parameters())
    assert parameter_count == 368, parameter_count
    assert model.readout.weight is model.token_embedding.weight

    train_inputs, train_targets, train_counts = make_dataset(config.train_size, config.train_seed)
    validation_inputs, validation_targets, validation_counts = make_dataset(
        config.validation_size, config.validation_seed
    )
    if one_sequence:
        train_inputs = torch.tensor([SEQUENCES[0][:-1]], dtype=torch.long)
        train_targets = torch.tensor([SEQUENCES[0][1:]], dtype=torch.long)
        train_counts = [1, 0, 0, 0]

    optimizer = torch.optim.AdamW(
        model.parameters(), lr=config.learning_rate, weight_decay=0.0
    )
    initial_train = evaluate(model, train_inputs, train_targets)
    initial_validation = evaluate(model, validation_inputs, validation_targets)
    batch_generator = torch.Generator().manual_seed(config.train_seed + 1)
    log_steps = {0, 1, 2, 5, 10, 25, 50, 100, 250, 500, config.steps}
    history = []

    for step in range(1, config.steps + 1):
        model.train()
        if one_sequence:
            batch_inputs, batch_targets = train_inputs, train_targets
        else:
            indices = torch.randint(
                len(train_inputs), (config.batch_size,), generator=batch_generator
            )
            batch_inputs = train_inputs[indices]
            batch_targets = train_targets[indices]
        optimizer.zero_grad(set_to_none=True)
        logits = model(batch_inputs)
        loss = token_loss(logits, batch_targets)
        loss.backward()
        before_clip = gradient_norm(model)
        if config.clip_norm is not None:
            torch.nn.utils.clip_grad_norm_(model.parameters(), config.clip_norm)
        after_clip = gradient_norm(model)
        optimizer.step()
        if step in log_steps:
            history.append(
                {
                    "step": step,
                    "batch_loss": float(loss.detach()),
                    "train_loss": evaluate(model, train_inputs, train_targets),
                    "validation_loss": evaluate(model, validation_inputs, validation_targets),
                    "gradient_norm_before_clip": before_clip,
                    "gradient_norm_after_clip": after_clip,
                }
            )

    model.eval()
    probes = torch.tensor([sequence[:-1] for sequence in SEQUENCES], dtype=torch.long)
    with torch.no_grad():
        probabilities = model(probes).softmax(dim=-1)
    selected_probes = {
        "after_bos": probabilities[0, 0].tolist(),
        "after_bos_a_b": probabilities[0, 2].tolist(),
        "after_bos_x_y": probabilities[2, 2].tolist(),
    }
    final_train_loss = evaluate(model, train_inputs, train_targets)
    final_validation_loss = evaluate(model, validation_inputs, validation_targets)

    saved = {
        "model": copy.deepcopy(model.state_dict()),
        "optimizer": copy.deepcopy(optimizer.state_dict()),
        "step": config.steps,
        "config": asdict(config),
        "vocabulary": VOCAB,
        "train_counts": train_counts,
        "validation_counts": validation_counts,
        "batch_generator_state": batch_generator.get_state(),
    }
    with tempfile.TemporaryDirectory() as directory:
        checkpoint_path = Path(directory) / "tiny-decoder.pt"
        torch.save(saved, checkpoint_path)
        loaded = torch.load(checkpoint_path, weights_only=True)
    torch.manual_seed(config.model_seed + 1000)
    restored = TinyDecoder(config)
    restored.load_state_dict(loaded["model"])
    assert torch.equal(model(probes), restored(probes))

    restored_optimizer = torch.optim.AdamW(
        restored.parameters(), lr=config.learning_rate, weight_decay=0.0
    )
    restored_optimizer.load_state_dict(loaded["optimizer"])
    restored_batch_generator = torch.Generator()
    restored_batch_generator.set_state(loaded["batch_generator_state"])

    def resume_one_step(
        resumed_model: TinyDecoder,
        resumed_optimizer: torch.optim.Optimizer,
        resumed_generator: torch.Generator,
    ) -> None:
        resumed_model.train()
        if one_sequence:
            next_inputs, next_targets = train_inputs, train_targets
        else:
            next_indices = torch.randint(
                len(train_inputs), (config.batch_size,), generator=resumed_generator
            )
            next_inputs = train_inputs[next_indices]
            next_targets = train_targets[next_indices]
        resumed_optimizer.zero_grad(set_to_none=True)
        next_loss = token_loss(resumed_model(next_inputs), next_targets)
        next_loss.backward()
        if config.clip_norm is not None:
            torch.nn.utils.clip_grad_norm_(resumed_model.parameters(), config.clip_norm)
        resumed_optimizer.step()

    resume_one_step(model, optimizer, batch_generator)
    resume_one_step(restored, restored_optimizer, restored_batch_generator)
    exact_resume = all(
        torch.equal(left, right)
        for left, right in zip(model.state_dict().values(), restored.state_dict().values())
    )
    assert exact_resume

    result = {
        "one_sequence": one_sequence,
        "torch_version": torch.__version__,
        "parameter_count": parameter_count,
        "config": asdict(config),
        "train_counts": train_counts,
        "validation_counts": validation_counts,
        "initial_train_loss": initial_train,
        "initial_validation_loss": initial_validation,
        "final_train_loss": final_train_loss,
        "final_validation_loss": final_validation_loss,
        "history": history,
        "probes": selected_probes,
        "checkpoint_reload_exact": True,
        "checkpoint_next_step_exact": exact_resume,
        "checkpoint_serialization_tested": True,
    }
    if return_state:
        return result, copy.deepcopy(saved["model"])
    return result


def main() -> None:
    base = Config()
    results = {
        "entropy_floor_nats_per_token": math.log(2) / 2,
        "one_sequence": train(
            Config(
                init_std=base.init_std,
                model_seed=base.model_seed,
                train_size=1,
                validation_size=base.validation_size,
                batch_size=1,
                steps=500,
                learning_rate=0.01,
                clip_norm=1.0,
            ),
            one_sequence=True,
        ),
        "corpus": train(base),
    }
    print(json.dumps(results, indent=2))


if __name__ == "__main__":
    main()
