"""Vectorize three ticket classifiers and inspect stable softmax."""

import json

import matplotlib.pyplot as plt
import numpy as np

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


CLASSES = ["security", "billing", "product"]


def stable_softmax(logits):
    shifted = logits - logits.max(axis=-1, keepdims=True)
    exponentials = np.exp(shifted)
    return exponentials / exponentials.sum(axis=-1, keepdims=True)


def main():
    x = np.array([[1.0, 0.0], [1.0, 1.0], [0.0, 1.0]])
    weights = np.array([[2.0, 0.5, -0.5], [0.0, -0.25, 1.0]])
    bias = np.array([0.0, 0.0, -0.5])
    labels = np.array([0, 0, 2])

    logits = x @ weights + bias
    probabilities = stable_softmax(logits)
    loss = float(-np.log(probabilities[np.arange(len(labels)), labels]).mean())
    dlogits = probabilities.copy()
    dlogits[np.arange(len(labels)), labels] -= 1
    dlogits /= len(labels)
    weight_gradient = x.T @ dlogits
    bias_gradient = dlogits.sum(axis=0)

    result = {
        "X_shape": list(x.shape),
        "W_shape": list(weights.shape),
        "b_shape": list(bias.shape),
        "Z_shape": list(logits.shape),
        "first_ticket_logits": logits[0].tolist(),
        "first_ticket_probabilities": probabilities[0].tolist(),
        "batch_cross_entropy": loss,
        "weight_gradient": weight_gradient.tolist(),
        "bias_gradient": bias_gradient.tolist(),
    }
    result_path("03_batch_and_softmax.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})
    colors = [BLUE, CYAN, GREEN]
    left.bar(CLASSES, logits[0], color=colors)
    left.axhline(0, color="#111317", linewidth=1)
    left.set(title="Raw output-layer scores", ylabel="logit", ylim=(-1.25, 2.3))
    left.grid(axis="y")
    for index, value in enumerate(logits[0]):
        left.text(index, value + (0.08 if value >= 0 else -0.18), f"{value:.1f}", ha="center", fontweight="bold")

    right.bar(CLASSES, probabilities[0], color=colors)
    right.set(title="Stable softmax distribution", ylabel="probability", ylim=(0, 1))
    right.grid(axis="y")
    for index, value in enumerate(probabilities[0]):
        right.text(index, value + 0.035, f"{value:.3f}", ha="center", fontweight="bold")
    fig.suptitle("One ticket: [2.0, 0.5, −1.0] becomes one normalized prediction", y=0.98, fontsize=20, fontweight="bold")
    fig.subplots_adjust(top=0.77)
    path = save_chart(fig, "03_batch_and_softmax.png")

    print("BATCH + STABLE SOFTMAX")
    print(f"X{x.shape} @ W{weights.shape} + b{bias.shape} -> Z{logits.shape}")
    print(f"first logits = {logits[0]}")
    print(f"shifted logits = {logits[0] - logits[0].max()}")
    print(f"first probabilities = {probabilities[0]}; sum = {probabilities[0].sum():.6f}")
    print(f"batch cross-entropy = {loss:.6f}")
    print(f"dW shape = {weight_gradient.shape}; db shape = {bias_gradient.shape}")
    print(f"chart = {path}")


if __name__ == "__main__":
    main()
