"""Plot a batch loss function, its local derivative, and one real update."""

import json

import matplotlib.pyplot as plt
import numpy as np

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


EXAMPLES = np.array(
    [
        [-2.0, 0.0],
        [-1.0, 0.0],
        [-0.25, 1.0],
        [0.25, 0.0],
        [1.0, 1.0],
        [2.0, 1.0],
    ]
)


def sigmoid(z):
    return np.where(z >= 0, 1.0 / (1.0 + np.exp(-z)), np.exp(z) / (1.0 + np.exp(z)))


def batch_loss(weight):
    x, y = EXAMPLES[:, 0], EXAMPLES[:, 1]
    z = weight * x
    per_example = np.maximum(z, 0) + np.log1p(np.exp(-np.abs(z))) - y * z
    return float(per_example.mean())


def batch_gradient(weight):
    x, y = EXAMPLES[:, 0], EXAMPLES[:, 1]
    return float(((sigmoid(weight * x) - y) * x).mean())


def main():
    weight = -2.0
    learning_rate = 2.0
    loss = batch_loss(weight)
    gradient = batch_gradient(weight)
    next_weight = weight - learning_rate * gradient
    next_loss = batch_loss(next_weight)

    grid = np.linspace(-5, 5, 801)
    losses = np.array([batch_loss(value) for value in grid])
    optimum_index = int(np.argmin(losses))
    optimum_weight = float(grid[optimum_index])
    optimum_loss = float(losses[optimum_index])
    tangent = loss + gradient * (grid - weight)

    result = {
        "weight_before": weight,
        "loss_before": loss,
        "derivative": gradient,
        "learning_rate": learning_rate,
        "weight_after": next_weight,
        "loss_after": next_loss,
        "plotted_optimum_weight": optimum_weight,
        "plotted_optimum_loss": optimum_loss,
    }
    result_path("02_loss_and_gradient.json").write_text(json.dumps(result, indent=2) + "\n")

    setup_chart()
    fig, ax = plt.subplots(figsize=(10.5, 5.6))
    ax.plot(grid, losses, color=BLUE, linewidth=4, label="actual batch loss L(w)")
    ax.plot(grid, tangent, color=RED, linewidth=2.4, linestyle="--", label="local tangent at w = −2")
    ax.scatter([weight], [loss], s=130, color=RED, zorder=5)
    ax.scatter([next_weight], [next_loss], s=150, facecolor="white", edgecolor=CYAN, linewidth=4, zorder=6)
    ax.scatter([optimum_weight], [optimum_loss], s=80, color=BLUE, zorder=5)
    ax.annotate(f"start\nw={weight:.2f}, L={loss:.3f}\nslope={gradient:.3f}", (weight, loss), xytext=(-108, 48), textcoords="offset points", arrowprops={"arrowstyle": "->", "color": RED}, fontsize=12)
    ax.annotate(f"one SGD step\nw={next_weight:.3f}, L={next_loss:.3f}", (next_weight, next_loss), xytext=(28, 55), textcoords="offset points", arrowprops={"arrowstyle": "->", "color": CYAN}, fontsize=12)
    ax.annotate(f"plotted minimum\nw≈{optimum_weight:.2f}", (optimum_weight, optimum_loss), xytext=(36, -54), textcoords="offset points", arrowprops={"arrowstyle": "->", "color": BLUE}, fontsize=12)
    ax.set(title="The derivative is a local rate of change—not the distance to the optimum", xlabel="One neuron weight, w", ylabel="Average binary cross-entropy", xlim=(-5, 5), ylim=(0, min(3.1, float(losses.max()) + 0.15)))
    ax.grid(True)
    ax.legend(loc="upper right", frameon=False)
    path = save_chart(fig, "02_loss_and_gradient.png")

    print("LOSS CURVE: RATE OF CHANGE VS. OPTIMUM")
    print(f"start: w={weight:.6f}, L={loss:.6f}, dL/dw={gradient:.6f}")
    print(f"update: w - eta*gradient = {weight:.6f} - {learning_rate:.1f}*({gradient:.6f}) = {next_weight:.6f}")
    print(f"recomputed loss: {next_loss:.6f}")
    print(f"plotted optimum: w≈{optimum_weight:.6f}, L≈{optimum_loss:.6f}")
    print(f"chart = {path}")


if __name__ == "__main__":
    main()
