"""
Lecture 07 - Gradient Descent & Stochastic Gradient Descent

Reproduces the exact worked numerical examples from the lecture page:
  Convergent run: L(w) = (w-3)^2, w0 = 0, eta = 0.1 -> converges toward w=3
  Divergent run:  same loss, eta = 1.1 -> oscillates and diverges

Also includes a small mini-batch gradient descent demo on a toy linear
regression dataset, illustrating the "practical default" discussed in the
lecture (batch_size = 32).

Run: python lecture-07-gradient-descent.py
"""
import numpy as np


def L(w):
    return (w - 3) ** 2


def grad_L(w):
    return 2 * (w - 3)


def gradient_descent(w0, lr, steps):
    w = w0
    history = [w]
    for _ in range(steps):
        w = w - lr * grad_L(w)
        history.append(w)
    return history


if __name__ == "__main__":
    # ---- convergent run: matches the lecture's worked example ----
    conv = gradient_descent(w0=0.0, lr=0.1, steps=6)
    print("Convergent (eta=0.1):", [round(w, 4) for w in conv])

    # ---- divergent run: learning rate too large ----
    div = gradient_descent(w0=0.0, lr=1.1, steps=4)
    print("Divergent  (eta=1.1):", [round(w, 4) for w in div])

    # ---- mini-batch gradient descent on a toy linear regression ----
    print("\nMini-batch gradient descent on a toy linear regression:")
    rng = np.random.default_rng(0)
    X = rng.uniform(-1, 1, size=(200, 1))
    true_w, true_b = 2.5, -0.7
    y = true_w * X[:, 0] + true_b + rng.normal(0, 0.05, size=200)

    w, b, lr, batch_size, epochs = 0.0, 0.0, 0.1, 32, 100
    n = len(X)
    for epoch in range(epochs):
        idx = rng.permutation(n)
        for start in range(0, n, batch_size):
            batch = idx[start:start + batch_size]
            xb, yb = X[batch, 0], y[batch]
            pred = w * xb + b
            err = pred - yb
            dw = 2 * np.mean(err * xb)
            db = 2 * np.mean(err)
            w -= lr * dw
            b -= lr * db
        if epoch % 20 == 0:
            mse = np.mean((w * X[:, 0] + b - y) ** 2)
            print(f"  epoch {epoch:3d}: w={w:.3f} b={b:.3f} mse={mse:.4f}")

    print(f"\nLearned w={w:.3f}, b={b:.3f}  (true w={true_w}, b={true_b})")
