"""
Gradient Descent Animation — Capítulo 3 (Visão Computacional, HM Lab)

Gera um GIF animado mostrando a descida do gradiente sobre a parábola
L(w) = (w - 3)^2, partindo de w0 = 0 com taxa de aprendizado eta = 0.2.

A cada iteração, o GIF mostra:
  - a curva da perda L(w)
  - o ponto atual (w, L(w)) em vermelho
  - a reta tangente no ponto (cuja inclinação = derivada = gradiente)
  - o rastro das iterações anteriores
  - um painel com w, L(w) e dL/dw da iteração corrente

Uso:
    pip install numpy matplotlib pillow
    python gradient_descent_animation.py

Saída:
    gradient_descent_animation.gif (na mesma pasta do script)
"""

import numpy as np
import matplotlib.pyplot as plt
from matplotlib.animation import FuncAnimation, PillowWriter

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

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

# ---------- hiperparâmetros ----------
W0 = 0.0          # ponto inicial
ETA = 0.2         # taxa de aprendizado
N_STEPS = 15      # número de iterações
FPS = 2           # frames por segundo do GIF (mais baixo = mais lento)

# ---------- pré-computa toda a trajetória ----------
ws = [W0]
for _ in range(N_STEPS):
    ws.append(ws[-1] - ETA * dL(ws[-1]))
ws = np.array(ws)

# grade da curva para plotar
xs = np.linspace(-1, 7, 300)
ys = L(xs)

# ---------- figura ----------
fig, ax = plt.subplots(figsize=(8, 5))
ax.plot(xs, ys, color="#4a90e2", lw=2.2, label=r"$L(w) = (w-3)^2$")
ax.plot(3, 0, "*", color="#27ae60", ms=18, zorder=5, label=r"mínimo $w^\ast = 3$")
ax.set_xlabel(r"$w$")
ax.set_ylabel(r"$L(w)$")
ax.grid(alpha=0.3)
ax.set_xlim(-1.2, 7.2)
ax.set_ylim(-1.5, ys.max() + 1)

# elementos animados (atualizados a cada frame)
trail,   = ax.plot([], [], "o-", color="#e74c3c", ms=5, lw=1.0, alpha=0.55,
                   label="trajetória")
current, = ax.plot([], [], "o", color="#e74c3c", ms=11, zorder=6,
                   markeredgecolor="white", markeredgewidth=1.2)
tangent, = ax.plot([], [], "--", color="#e74c3c", lw=2.0, alpha=0.85,
                   label="reta tangente")
info = ax.text(0.02, 0.96, "", transform=ax.transAxes, va="top", ha="left",
               fontsize=11, family="monospace",
               bbox=dict(boxstyle="round,pad=0.45", fc="white", ec="#bbb", alpha=0.9))

ax.legend(loc="upper right", fontsize=9.5, framealpha=0.92)
ax.set_title("Descida do gradiente em $L(w)=(w-3)^2$  —  $\\eta = 0.2$, $w_0 = 0$")

def init():
    trail.set_data([], [])
    current.set_data([], [])
    tangent.set_data([], [])
    info.set_text("")
    return trail, current, tangent, info

def update(i):
    w_now = ws[i]
    grad = dL(w_now)

    # rastro: pontos visitados até aqui
    trail.set_data(ws[:i+1], L(ws[:i+1]))

    # ponto atual
    current.set_data([w_now], [L(w_now)])

    # reta tangente local: y = L(w_now) + grad * (x - w_now), num pequeno intervalo
    half = 0.9
    x_tan = np.array([w_now - half, w_now + half])
    y_tan = L(w_now) + grad * (x_tan - w_now)
    tangent.set_data(x_tan, y_tan)

    info.set_text(
        f"iteração: {i:>2d}\n"
        f"w        = {w_now:+.4f}\n"
        f"L(w)     = {L(w_now):.4f}\n"
        f"dL/dw    = {grad:+.4f}\n"
        f"próximo passo: w - eta*dL/dw"
    )
    return trail, current, tangent, info

anim = FuncAnimation(fig, update, frames=len(ws),
                     init_func=init, blit=True, interval=1000 // FPS)

OUT = "gradient_descent_animation.gif"
anim.save(OUT, writer=PillowWriter(fps=FPS))
print(f"GIF salvo em: {OUT}")
plt.close(fig)
