"""Train a tiny embedding plus softmax language model with manual NumPy gradients."""

import json

import matplotlib.pyplot as plt
import numpy as np
from embedding_model import count_model, train_model, predict

from course_utils import BLUE, CYAN, GREEN, result_path, save_chart, setup_chart

def main():
    ngram = count_model()
    fitted = train_model(ngram.CORPUS)
    model = fitted['trained']
    vocabulary, output_vocabulary = model['vocabulary'], model['outputs']
    embeddings = np.array(model['embeddings'])
    output_weights = np.array(model['weights'])
    contexts = np.array([[vocabulary.index(token) for token in context] for context, _ in fitted['rows']])
    losses = fitted['losses']
    query = predict(model, ['bug', 'is'])
    ranked = sorted(zip(output_vocabulary, query['probabilities']), key=lambda item: -item[1])

    result = {
        "vocabulary": vocabulary,
        "output_vocabulary": output_vocabulary,
        "embedding_shape": list(embeddings.shape),
        "output_weight_shape": list(output_weights.shape),
        "training_examples": len(contexts),
        "epochs": len(losses),
        "loss_start": losses[0],
        "loss_end": losses[-1],
        "query_context": ["bug", "is"],
        "top_probabilities": ranked[:5],
    }
    result_path("05_neural_ngram.json").write_text(json.dumps(result, indent=2) + "\n")

    setup_chart()
    fig, (left, right) = plt.subplots(1, 2, figsize=(11.2, 5.4), gridspec_kw={"wspace": 0.34})
    left.plot(np.arange(1, len(losses) + 1), losses, color=BLUE, linewidth=3)
    left.set(title="Training loss", xlabel="full-batch SGD step", ylabel="cross-entropy")
    left.grid(True)
    top = ranked[:3]
    right.bar([item[0] for item in top], [item[1] for item in top], color=[BLUE, CYAN, GREEN])
    right.set(title='P(next | “bug is”)', ylabel="probability", ylim=(0, 0.62))
    right.grid(axis="y")
    for index, (_, value) in enumerate(top):
        right.text(index, value + 0.025, f"{value:.3f}", ha="center", fontweight="bold")
    fig.suptitle("A learned representation replaces the count lookup—not the objective", y=0.98, fontsize=20, fontweight="bold")
    fig.subplots_adjust(top=0.77)
    path = save_chart(fig, "05_neural_ngram.png")

    print("NEURAL TRIGRAM: EMBEDDINGS + LINEAR LM HEAD")
    print(f"contexts = {contexts.shape}; embeddings = {embeddings.shape}; output weights = {output_weights.shape}")
    print(f"loss: {losses[0]:.6f} -> {losses[-1]:.6f}")
    print(f"top P(next | ('bug', 'is')) = {ranked[:3]}")
    print(f"chart = {path}")


if __name__ == "__main__":
    main()
