#!/usr/bin/env python3
"""Train and evaluate a tiny contrastive embedding model on CPU.

The data is deliberately artificial. Words are arranged into five topic groups,
and positive pairs come from the same group. Therefore, learned neighbours show
that the model recovered this toy co-occurrence rule; they are not claims about
the meanings of words in natural language.
"""

from __future__ import annotations

import argparse

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


torch.set_num_threads(1)


TOPIC_WORDS = {
    "animals": ("cat", "dog", "kitten", "puppy"),
    "fruit": ("apple", "pear", "peach", "plum"),
    "music": ("piano", "violin", "drum", "flute"),
    "weather": ("rain", "snow", "wind", "cloud"),
    "travel": ("train", "bus", "plane", "boat"),
}
VOCABULARY = tuple(word for words in TOPIC_WORDS.values() for word in words)
WORD_TO_ID = {word: index for index, word in enumerate(VOCABULARY)}
WORD_TO_TOPIC = {
    word: topic for topic, words in TOPIC_WORDS.items() for word in words
}


def make_contrastive_pairs(
    pair_count: int,
    seed: int,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
    """Return balanced centre/context pairs and binary labels.

    A label of one means both words came from one topic. A label of zero means
    they came from different topics. Each tensor has shape ``(pair_count,)``.
    """
    if pair_count < 2 or pair_count % 2:
        raise ValueError("pair_count must be an even number of at least two")

    generator = torch.Generator(device="cpu").manual_seed(seed)
    groups = tuple(tuple(WORD_TO_ID[word] for word in words) for words in TOPIC_WORDS.values())
    centres: list[int] = []
    contexts: list[int] = []
    labels: list[float] = []

    for index in range(pair_count // 2):
        topic_index = index % len(groups)
        group = groups[topic_index]
        first = torch.randint(len(group), (1,), generator=generator).item()
        offset = torch.randint(1, len(group), (1,), generator=generator).item()
        centres.append(group[first])
        contexts.append(group[(first + offset) % len(group)])
        labels.append(1.0)

        other_offset = torch.randint(1, len(groups), (1,), generator=generator).item()
        other_group = groups[(topic_index + other_offset) % len(groups)]
        centres.append(group[first])
        contexts.append(
            other_group[torch.randint(len(other_group), (1,), generator=generator).item()]
        )
        labels.append(0.0)

    order = torch.randperm(pair_count, generator=generator)
    return (
        torch.tensor(centres, dtype=torch.long)[order],
        torch.tensor(contexts, dtype=torch.long)[order],
        torch.tensor(labels, dtype=torch.float32)[order],
    )


class ContrastiveEmbedding(nn.Module):
    """One shared lookup table trained by same-context comparisons."""

    def __init__(self, vocabulary_size: int, embedding_size: int = 8) -> None:
        super().__init__()
        self.embedding = nn.Embedding(vocabulary_size, embedding_size)

    def forward(self, centre_ids: torch.Tensor, context_ids: torch.Tensor) -> torch.Tensor:
        centre = F.normalize(self.embedding(centre_ids), dim=-1)
        context = F.normalize(self.embedding(context_ids), dim=-1)
        return 5.0 * (centre * context).sum(dim=-1)


def pair_accuracy(
    model: ContrastiveEmbedding,
    centres: torch.Tensor,
    contexts: torch.Tensor,
    labels: torch.Tensor,
) -> tuple[float, float]:
    """Return held-out binary cross-entropy and accuracy."""
    model.eval()
    with torch.no_grad():
        logits = model(centres, contexts)
        loss = F.binary_cross_entropy_with_logits(logits, labels).item()
        accuracy = ((logits >= 0) == labels.bool()).float().mean().item()
    return loss, accuracy


def retrieval_precision(model: ContrastiveEmbedding, neighbours: int = 3) -> float:
    """Measure how often nearest neighbours share the toy topic label."""
    model.eval()
    with torch.no_grad():
        vectors = F.normalize(model.embedding.weight, dim=-1)
        similarities = vectors @ vectors.T
        similarities.fill_diagonal_(-float("inf"))
        nearest = similarities.topk(neighbours, dim=1).indices

    correct = 0
    for word_id, neighbour_ids in enumerate(nearest.tolist()):
        topic = WORD_TO_TOPIC[VOCABULARY[word_id]]
        correct += sum(
            WORD_TO_TOPIC[VOCABULARY[neighbour_id]] == topic
            for neighbour_id in neighbour_ids
        )
    return correct / (len(VOCABULARY) * neighbours)


def nearest_words(
    model: ContrastiveEmbedding,
    query: str,
    neighbours: int = 3,
) -> list[str]:
    """Return nearest words by cosine similarity, excluding the query."""
    query_id = WORD_TO_ID[query]
    with torch.no_grad():
        vectors = F.normalize(model.embedding.weight, dim=-1)
        scores = vectors @ vectors[query_id]
        scores[query_id] = -float("inf")
        ids = scores.topk(neighbours).indices.tolist()
    return [VOCABULARY[index] for index in ids]


def train_model(steps: int = 250) -> ContrastiveEmbedding:
    """Train with a fixed generated training set."""
    if steps < 1:
        raise ValueError("steps must be positive")
    torch.manual_seed(2026)
    train_centres, train_contexts, train_labels = make_contrastive_pairs(800, seed=10)
    model = ContrastiveEmbedding(len(VOCABULARY), embedding_size=8)
    optimizer = torch.optim.Adam(model.parameters(), lr=0.04)

    model.train()
    for _ in range(steps):
        optimizer.zero_grad()
        logits = model(train_centres, train_contexts)
        loss = F.binary_cross_entropy_with_logits(logits, train_labels)
        loss.backward()
        optimizer.step()
    return model


def run(steps: int = 250, smoke_test: bool = False) -> dict[str, float]:
    """Run training, held-out evaluation, baseline comparison, and retrieval."""
    effective_steps = min(steps, 3) if smoke_test else steps
    validation = make_contrastive_pairs(400, seed=20)
    model = train_model(steps=effective_steps)
    validation_loss, validation_accuracy = pair_accuracy(model, *validation)
    baseline_accuracy = max(validation[2].mean().item(), 1.0 - validation[2].mean().item())
    precision = retrieval_precision(model)

    print(f"vocabulary size: {len(VOCABULARY)}")
    print("lookup shape for 4 IDs:", tuple(model.embedding(torch.arange(4)).shape))
    print(f"majority baseline accuracy: {baseline_accuracy:.3f}")
    print(f"held-out pair loss: {validation_loss:.4f}")
    print(f"held-out pair accuracy: {validation_accuracy:.3f}")
    print(f"same-topic precision@3: {precision:.3f}")
    print("nearest to cat:", ", ".join(nearest_words(model, "cat")))
    print("nearest to piano:", ", ".join(nearest_words(model, "piano")))
    if smoke_test:
        print("smoke test: learning thresholds skipped")
    else:
        assert validation_accuracy > baseline_accuracy + 0.30
        assert precision > 0.85

    return {
        "baseline_accuracy": baseline_accuracy,
        "validation_loss": validation_loss,
        "validation_accuracy": validation_accuracy,
        "retrieval_precision": precision,
    }


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--steps", type=int, default=250)
    parser.add_argument(
        "--smoke-test",
        action="store_true",
        help="run at most three optimizer steps and skip learning thresholds",
    )
    return parser.parse_args()


if __name__ == "__main__":
    arguments = parse_args()
    run(steps=arguments.steps, smoke_test=arguments.smoke_test)
