"""Train a tiny transformer on copying integer sequences; no downloads or GPU.

Keep this file next to torch_core.py. Run: python train_copy.py --steps 800
The train/evaluation sequences are disjoint; evaluation uses greedy generation.
"""
from __future__ import annotations

from typing import TypedDict
from torch import Tensor
import argparse
import itertools
import json
import time
import torch
from torch.nn import functional as F
from torch_core import TinyTransformer, generate

class GenerationExample(TypedDict):
    source: list[int]
    expected: list[int]
    generated: list[int]


class TrainingResult(TypedDict):
    seed: int
    steps: int
    train_sequences: int
    held_out_sequences: int
    initial_train_loss: float
    held_out_loss: float
    greedy_exact_match: float
    seconds: float
    numpy_torch_note: str
    examples: list[GenerationExample]


# region data
def dataset(seed: int = 7) -> tuple[list[list[int]], list[list[int]]]:
    sequences = [list(items) for length in [1, 2, 3]
                 for items in itertools.product(range(4, 12), repeat=length)]
    generator = torch.Generator().manual_seed(seed)
    order = torch.randperm(len(sequences), generator=generator).tolist()
    cut = int(0.8 * len(order))
    return ([sequences[i] for i in order[:cut]],
            [sequences[i] for i in order[cut:]])


def batch(sequences: list[list[int]]) -> tuple[Tensor, Tensor, Tensor]:
    source = torch.zeros(len(sequences), 3, dtype=torch.long)
    decoder_input = torch.zeros(len(sequences), 4, dtype=torch.long)
    target = torch.zeros(len(sequences), 4, dtype=torch.long)
    for i, tokens in enumerate(sequences):
        source[i, :len(tokens)] = torch.tensor(tokens)
        decoder_input[i, :len(tokens) + 1] = torch.tensor([1] + tokens)
        target[i, :len(tokens) + 1] = torch.tensor(tokens + [2])
    return source, decoder_input, target
# endregion data

# region loss
def training_loss(model: TinyTransformer, source: Tensor, decoder_input: Tensor, target: Tensor) -> Tensor:
    logits = model(source, decoder_input)  # (B, T, vocabulary)
    return F.cross_entropy(
        logits.reshape(-1, logits.shape[-1]),
        target.reshape(-1),
        ignore_index=0,                   # Padding is not a prediction target.
    )                                    # Pass raw logits; do not softmax first.
# endregion loss

# region train
def train(steps: int = 800, seed: int = 7) -> tuple[TinyTransformer, TrainingResult]:
    torch.set_num_threads(1)
    torch.manual_seed(seed)
    train_sequences, held_out = dataset(seed)
    model = TinyTransformer()
    optimizer = torch.optim.Adam(model.parameters(), lr=1e-3)
    generator = torch.Generator().manual_seed(seed + 1)
    initial_source, initial_input, initial_target = batch(train_sequences[:32])
    initial_loss = training_loss(model, initial_source, initial_input, initial_target).item()
    start = time.perf_counter()
    for _ in range(steps):
        indices = torch.randint(len(train_sequences), (32,), generator=generator)
        source, decoder_input, target = batch([train_sequences[i] for i in indices.tolist()])
        optimizer.zero_grad(set_to_none=True)
        loss = training_loss(model, source, decoder_input, target)
        loss.backward()
        torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
        optimizer.step()

    source, decoder_input, target = batch(held_out)
    model.eval()
    with torch.no_grad():
        held_out_loss = training_loss(model, source, decoder_input, target).item()
        generated = generate(model, source)
    generated = F.pad(generated, (0, target.shape[1] - generated.shape[1]))
    exact_match = (generated == target).all(dim=1).float().mean().item()
    result: TrainingResult = {
        "seed": seed, "steps": steps, "train_sequences": len(train_sequences),
        "held_out_sequences": len(held_out), "initial_train_loss": round(initial_loss, 6),
        "held_out_loss": round(held_out_loss, 6), "greedy_exact_match": round(exact_match, 6),
        "seconds": round(time.perf_counter() - start, 2),
        "numpy_torch_note": "NumPy demonstrates forward passes; PyTorch trains the model.",
        "examples": [
            {"source": source[i].tolist(), "expected": target[i].tolist(),
             "generated": generated[i].tolist()} for i in range(3)
        ],
    }
    return model, result
# endregion train

if __name__ == "__main__":
    parser = argparse.ArgumentParser()
    parser.add_argument("--steps", type=int, default=800)
    parser.add_argument("--seed", type=int, default=7)
    args = parser.parse_args()
    if args.steps < 1:
        parser.error("--steps must be positive")
    _, result = train(args.steps, args.seed)
    print(json.dumps(result, indent=2))
