"""A small decoder-only transformer. CPU, character tokens, no pretrained weights."""
from dataclasses import asdict, dataclass
import hashlib
from pathlib import Path

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

DATA_SHA256 = '86c4e6aa9db7c042ec79f339dcb96d42b0075e16b8fc2e86bf0ca57e2dc565ed'


@dataclass
class Config:
    vocab_size: int = 65
    context_length: int = 64
    width: int = 64
    heads: int = 4
    layers: int = 2


class Attention(nn.Module):
    """Mix earlier information; keep each head's weights available for inspection."""
    def __init__(self, config):
        super().__init__()
        assert config.width % config.heads == 0
        self.heads = config.heads
        self.qkv = nn.Linear(config.width, 3 * config.width)
        self.output = nn.Linear(config.width, config.width)
        self.register_buffer('allowed', torch.tril(torch.ones(
            config.context_length, config.context_length, dtype=torch.bool)))

    def forward(self, x):
        batch, positions, width = x.shape
        query, key, value = self.qkv(x).chunk(3, dim=-1)
        # Give each head its own slice of features.
        query, key, value = [part.view(batch, positions, self.heads,
                                       width // self.heads).transpose(1, 2)
                             for part in (query, key, value)]
        scores = query @ key.transpose(-2, -1) / (width // self.heads) ** 0.5
        scores = scores.masked_fill(~self.allowed[:positions, :positions], float('-inf'))
        weights = F.softmax(scores, dim=-1)
        mixed = weights @ value
        mixed = mixed.transpose(1, 2).contiguous().view(batch, positions, width)
        return self.output(mixed), weights


class Block(nn.Module):
    def __init__(self, config):
        super().__init__()
        self.before_attention = nn.LayerNorm(config.width)
        self.attention = Attention(config)
        self.before_mlp = nn.LayerNorm(config.width)
        self.mlp = nn.Sequential(nn.Linear(config.width, 4 * config.width),
                                 nn.GELU(), nn.Linear(4 * config.width, config.width))

    def forward(self, x):
        # Normalization controls scale. Addition preserves a path for existing information.
        mixture, weights = self.attention(self.before_attention(x))
        x = x + mixture
        x = x + self.mlp(self.before_mlp(x))
        return x, weights


class TinyTransformer(nn.Module):
    def __init__(self, config=Config()):
        super().__init__()
        self.config = config
        self.token_embedding = nn.Embedding(config.vocab_size, config.width)
        self.position_embedding = nn.Embedding(config.context_length, config.width)
        self.blocks = nn.ModuleList([Block(config) for _ in range(config.layers)])
        self.final_norm = nn.LayerNorm(config.width)
        self.output = nn.Linear(config.width, config.vocab_size)
        self.apply(self._initialize)

    @staticmethod
    def _initialize(module):
        if isinstance(module, (nn.Embedding, nn.Linear)):
            nn.init.normal_(module.weight, mean=0, std=0.02)
        if isinstance(module, nn.Linear) and module.bias is not None:
            nn.init.zeros_(module.bias)

    def forward(self, token_ids, return_attention=False):
        if token_ids.ndim != 2 or not 0 < token_ids.shape[1] <= self.config.context_length:
            raise ValueError('Expected [batch, positions] with 1..context_length positions.')
        positions = torch.arange(token_ids.shape[1], device=token_ids.device)
        x = self.token_embedding(token_ids) + self.position_embedding(positions)
        attention = []
        for block in self.blocks:
            x, weights = block(x)
            if return_attention:
                attention.append(weights)
        logits = self.output(self.final_norm(x))
        return (logits, attention) if return_attention else logits


class CharacterTokenizer:
    def __init__(self, characters):
        self.characters = list(characters)
        self.to_id = {char: i for i, char in enumerate(self.characters)}

    def encode(self, text):
        unknown = set(text) - self.to_id.keys()
        if unknown:
            raise ValueError(f'Characters outside this lab vocabulary: {sorted(unknown)!r}')
        return [self.to_id[char] for char in text]

    def decode(self, ids):
        return ''.join(self.characters[int(i)] for i in ids)


def load_data(path=None):
    path = Path(path) if path else Path(__file__).with_name('input.txt')
    raw = path.read_bytes()
    if hashlib.sha256(raw).hexdigest() != DATA_SHA256:
        raise ValueError('input.txt does not match the published Tiny Shakespeare corpus.')
    text = raw.decode('utf-8')
    tokenizer = CharacterTokenizer(sorted(set(text)))
    split = int(0.9 * len(text))
    # Split the continuous stream BEFORE taking windows: no train/validation overlap.
    train = torch.tensor(tokenizer.encode(text[:split]), dtype=torch.long)
    validation = torch.tensor(tokenizer.encode(text[split:]), dtype=torch.long)
    return tokenizer, train, validation


def get_batch(stream, batch_size, context_length, generator):
    if len(stream) <= context_length:
        raise ValueError('A stream must supply context_length + 1 characters.')
    starts = torch.randint(len(stream) - context_length, (batch_size,), generator=generator)
    x = torch.stack([stream[start:start + context_length] for start in starts])
    y = torch.stack([stream[start + 1:start + context_length + 1] for start in starts])
    return x, y


def prediction_loss(logits, targets):
    # Cross-entropy consumes raw scores; it includes the needed normalization.
    return F.cross_entropy(logits.reshape(-1, logits.shape[-1]), targets.reshape(-1))


@torch.no_grad()
def generate(model, tokenizer, prompt='\nROMEO:\n', new_tokens=160, temperature=1.0, seed=2026):
    if not prompt:
        raise ValueError('This stream model needs a nonempty prompt; it has no BOS token.')
    if temperature <= 0 or new_tokens < 0:
        raise ValueError('Temperature must be positive; new_tokens must be nonnegative.')
    rng = torch.Generator(device='cpu').manual_seed(seed)
    ids = torch.tensor([tokenizer.encode(prompt)], dtype=torch.long)
    was_training = model.training
    model.eval()
    try:
        for _ in range(new_tokens):
            context = ids[:, -model.config.context_length:]
            logits = model(context)[:, -1, :]
            probabilities = F.softmax(logits / temperature, dim=-1)
            next_id = torch.multinomial(probabilities, 1, generator=rng)
            ids = torch.cat((ids, next_id), dim=1)
    finally:
        model.train(was_training)
    return tokenizer.decode(ids[0])


def save_checkpoint(path, model, tokenizer, step):
    torch.save({'config': asdict(model.config), 'characters': tokenizer.characters,
                'state_dict': model.state_dict(), 'step': step, 'data_sha256': DATA_SHA256}, path)


def load_checkpoint(path):
    saved = torch.load(path, map_location='cpu', weights_only=True)
    if saved['data_sha256'] != DATA_SHA256:
        raise ValueError('Checkpoint was trained on a different corpus.')
    model = TinyTransformer(Config(**saved['config']))
    model.load_state_dict(saved['state_dict'])
    model.eval()
    return model, CharacterTokenizer(saved['characters']), saved['step']
