"""Work one binary-neuron prediction and SGD update from end to end."""

import json

import matplotlib.pyplot as plt
import numpy as np

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


def sigmoid(z):
    return 1.0 / (1.0 + np.exp(-z))


def binary_cross_entropy(probability, label):
    probability = np.clip(probability, 1e-12, 1.0 - 1e-12)
    return -(label * np.log(probability) + (1 - label) * np.log(1 - probability))


def forward(x, weights, bias):
    score = float(x @ weights + bias)
    probability = float(sigmoid(score))
    return score, probability


def gradients(x, probability, label):
    score_error = probability - label
    return score_error * x, float(score_error)


def finite_difference_weight(x, weights, bias, label, index, epsilon=1e-6):
    plus = weights.copy()
    minus = weights.copy()
    plus[index] += epsilon
    minus[index] -= epsilon
    _, probability_plus = forward(x, plus, bias)
    _, probability_minus = forward(x, minus, bias)
    return float(
        (binary_cross_entropy(probability_plus, label) - binary_cross_entropy(probability_minus, label))
        / (2 * epsilon)
    )


def main():
    x = np.array([1.0, 0.0])
    weights = np.array([1.2, 0.7])
    bias = -0.5
    label = 1.0
    learning_rate = 0.1

    score, probability = forward(x, weights, bias)
    loss = float(binary_cross_entropy(probability, label))
    weight_gradient, bias_gradient = gradients(x, probability, label)
    numerical_weight_gradient = finite_difference_weight(x, weights, bias, label, index=0)
    next_weights = weights - learning_rate * weight_gradient
    next_bias = bias - learning_rate * bias_gradient
    next_score, next_probability = forward(x, next_weights, next_bias)
    next_loss = float(binary_cross_entropy(next_probability, label))

    result = {
        "x": x.tolist(),
        "weights_before": weights.tolist(),
        "bias_before": bias,
        "score_before": score,
        "probability_before": probability,
        "loss_before": loss,
        "score_error": probability - label,
        "weight_gradient": weight_gradient.tolist(),
        "finite_difference_weight_1": numerical_weight_gradient,
        "bias_gradient": bias_gradient,
        "learning_rate": learning_rate,
        "weights_after": next_weights.tolist(),
        "bias_after": next_bias,
        "score_after": next_score,
        "probability_after": next_probability,
        "loss_after": next_loss,
    }
    result_path("01_single_neuron.json").write_text(json.dumps(result, indent=2) + "\n")

    setup_chart()
    scores = np.linspace(-6, 6, 400)
    fig, ax = plt.subplots(figsize=(10.5, 5.6))
    ax.plot(scores, sigmoid(scores), color=BLUE, linewidth=4, label="sigmoid(z)")
    ax.axhline(0.5, color=INK, linewidth=1.2, linestyle="--", label="0.50 product threshold")
    ax.scatter([score], [probability], s=130, color=RED, zorder=4, label="before update")
    ax.scatter([next_score], [next_probability], s=150, facecolor="white", edgecolor=CYAN, linewidth=4, zorder=5, label="after update")
    ax.annotate(f"z={score:.3f}\np={probability:.3f}", (score, probability), xytext=(-92, 42), textcoords="offset points", arrowprops={"arrowstyle": "->", "color": RED}, fontsize=12)
    ax.annotate(f"z={next_score:.3f}\np={next_probability:.3f}", (next_score, next_probability), xytext=(24, -58), textcoords="offset points", arrowprops={"arrowstyle": "->", "color": CYAN}, fontsize=12)
    ax.set(title="One parameter update raises the positive-class probability", xlabel="Neuron score, z", ylabel="p(y=1)", xlim=(-6, 6), ylim=(-0.03, 1.03))
    ax.grid(True)
    ax.legend(loc="lower right", frameon=False)
    path = save_chart(fig, "01_single_neuron.png")

    print("ONE NEURON: FORWARD → LOSS → GRADIENT → UPDATE")
    print(f"z = {score:.6f}; p = {probability:.6f}; loss = {loss:.6f}")
    print(f"dL/dz = {probability - label:.6f}")
    print(f"dL/dw = {weight_gradient}; dL/db = {bias_gradient:.6f}")
    print(f"finite-difference dL/dw1 = {numerical_weight_gradient:.6f}")
    print(f"w_next = {next_weights}; b_next = {next_bias:.6f}")
    print(f"p_next = {next_probability:.6f}; loss_next = {next_loss:.6f}")
    print(f"chart = {path}")


if __name__ == "__main__":
    main()
