"""
Lecture 08 - Backpropagation from scratch (NumPy)

Reproduces the exact worked numerical example from the lecture page:
  x = [0.5, 0.8], target y = 1, learning rate = 0.1
  2 inputs -> 2 hidden units (ReLU) -> 1 output (sigmoid), binary cross-entropy loss

Run: python lecture-08-backprop.py
"""
import numpy as np


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


def relu(z):
    return np.maximum(0, z)


def relu_grad(z):
    return (z > 0).astype(float)


# ---- data & parameters (matches the worked example in the lecture) ----
X = np.array([[0.5, 0.8]])                      # (1, 2)
y = np.array([[1.0]])                            # target
W1 = np.array([[0.3, -0.1], [0.2, 0.4]])          # (2, 2)
b1 = np.array([[0.1, -0.2]])                      # (1, 2)
W2 = np.array([[0.5], [-0.3]])                    # (2, 1)
b2 = np.array([[0.05]])                           # (1, 1)
lr = 0.1


def forward(X, W1, b1, W2, b2):
    Z1 = X @ W1 + b1
    A = relu(Z1)
    Z2 = A @ W2 + b2
    Q = sigmoid(Z2)
    return Z1, A, Z2, Q


def backward(X, y, Z1, A, Q, W2):
    dZ2 = Q - y                       # sigmoid + BCE shortcut: dL/dZ2 = Q - y
    dW2 = A.T @ dZ2
    db2 = dZ2.sum(axis=0, keepdims=True)

    dA = dZ2 @ W2.T
    dZ1 = dA * relu_grad(Z1)
    dW1 = X.T @ dZ1
    db1 = dZ1.sum(axis=0, keepdims=True)
    return dW1, db1, dW2, db2


if __name__ == "__main__":
    Z1, A, Z2, Q = forward(X, W1, b1, W2, b2)
    L = -(y * np.log(Q) + (1 - y) * np.log(1 - Q))
    print("Forward pass")
    print("  Z1 =", Z1.ravel())
    print("  A  =", A.ravel())
    print("  Z2 =", Z2.ravel())
    print("  Q  =", Q.ravel())
    print("  Loss =", L.ravel())

    dW1, db1, dW2, db2 = backward(X, y, Z1, A, Q, W2)
    print("\nBackward pass (gradients)")
    print("  dW2 =", dW2.ravel())
    print("  db2 =", db2.ravel())
    print("  dW1 =\n", dW1)
    print("  db1 =", db1.ravel())

    # gradient descent update
    W1 -= lr * dW1
    b1 -= lr * db1
    W2 -= lr * dW2
    b2 -= lr * db2
    print("\nUpdated parameters (one gradient-descent step, eta=0.1)")
    print("  W1 =\n", W1)
    print("  b1 =", b1.ravel())
    print("  W2 =", W2.ravel())
    print("  b2 =", b2.ravel())

    # quick sanity check: loss should decrease after the update
    _, _, _, Q_new = forward(X, W1, b1, W2, b2)
    print("\nQ before update:", Q.ravel(), " Q after update:", Q_new.ravel(),
          " (should have moved closer to y=1)")

    # ---- mini training loop: watch the loss actually converge ----
    print("\nTraining for 200 steps on this single example...")
    W1 = np.array([[0.3, -0.1], [0.2, 0.4]])
    b1 = np.array([[0.1, -0.2]])
    W2 = np.array([[0.5], [-0.3]])
    b2 = np.array([[0.05]])
    for step in range(200):
        Z1, A, Z2, Q = forward(X, W1, b1, W2, b2)
        dW1, db1, dW2, db2 = backward(X, y, Z1, A, Q, W2)
        W1 -= lr * dW1; b1 -= lr * db1
        W2 -= lr * dW2; b2 -= lr * db2
        if step % 40 == 0:
            L = -(y * np.log(Q) + (1 - y) * np.log(1 - Q))
            print(f"  step {step:3d}: Q={Q.ravel()[0]:.4f}  loss={L.ravel()[0]:.4f}")
