"""
Animasi 2D: Dynamic Spiral Optimization
- Kandidat solusi (titik) tersebar acak
- Berputar (rotasi + kontraksi) mengelilingi titik terbaik
- Pusat spiral diperbarui dinamis saat ditemukan solusi lebih baik
- Konvergen ke solusi optimal
Output: MP4 1280x720 @ 30 fps
"""

import numpy as np
import matplotlib

matplotlib.use("Agg")
import matplotlib.pyplot as plt
from matplotlib.collections import LineCollection
import imageio.v2 as imageio

# ---------------- konfigurasi ----------------
FPS = 30
W, H = 12.8, 7.2  # inci -> 1280x720 @ 100 dpi
DPI = 100
N_POINTS = 40
N_ITER = 18
FRAMES_PER_ITER = 24
INTRO_FRAMES = 75
HILITE_FRAMES = 60
END_FRAMES = 105
TRAIL_LEN = 55
THETA = np.pi / 4  # sudut rotasi per iterasi

OUT = "/home/ikanx101/Documents/local-AI-agent-setup/Karya_Agent/IMW/resource/dynamic_spiral_optimization.mp4"

BG = "#0d1117"
ACCENT = "#37c8f0"
GOLD = "#ffcc33"
TXT = "#e8eef4"
DIM = "#8899aa"

OPT = np.array([1.8, -1.1])  # solusi optimal sebenarnya


def f(x, y):
    a = 0.6
    dx, dy = x - OPT[0], y - OPT[1]
    u = np.cos(a) * dx + np.sin(a) * dy
    v = -np.sin(a) * dx + np.cos(a) * dy
    return (u / 2.4) ** 2 + (v / 1.4) ** 2


def feval(p):
    return f(p[:, 0], p[:, 1])


def r_at(k):
    """kontraksi dinamis: mengecil seiring iterasi"""
    return max(0.72, 0.88 - 0.008 * k)


# ---------------- simulasi SPO ----------------
rng = np.random.default_rng(7)
pts0 = rng.uniform(-4.6, 4.6, size=(N_POINTS, 2))

# jauhkan titik awal dari optimum agar spiral terlihat jelas
mask = np.linalg.norm(pts0 - OPT, axis=1) < 1.2
pts0[mask] += np.array([2.5, 2.5])

iters = [pts0.copy()]
centers = []
rvals = []
pts = pts0.copy()
for k in range(N_ITER):
    fv = feval(pts)
    c = pts[np.argmin(fv)].copy()
    r = r_at(k)
    centers.append(c)
    rvals.append(r)
    ct, st = np.cos(THETA), np.sin(THETA)
    R = np.array([[ct, -st], [st, ct]])
    pts = c + r * (pts - c) @ R.T
    iters.append(pts.copy())

fbest_final = feval(iters[-1]).min()
print(f"sebaran akhir (std): {iters[-1].std():.4f}, f terbaik: {fbest_final:.5f}")

# ---------------- grid kontur ----------------
gx = np.linspace(-5.6, 5.6, 320)
gy = np.linspace(-5.6, 5.6, 320)
GX, GY = np.meshgrid(gx, gy)
GZ = f(GX, GY)
LEVELS = np.linspace(0, GZ.max() ** 0.5, 16) ** 2 + 0.02

# ---------------- helper ----------------

def smoothstep(t):
    return t * t * (3 - 2 * t)


def interp_positions(k, t):
    """posisi antar-iterasi: rotasi & kontraksi kontinu (busur spiral asli)"""
    p = iters[k]
    c = centers[k]
    r = rvals[k]
    ang = THETA * t
    rad = r ** t
    ct, st = np.cos(ang), np.sin(ang)
    R = np.array([[ct, -st], [st, ct]])
    return c + rad * (p - c) @ R.T


def caption_for(k):
    if k < 4:
        return "Rotasi + kontraksi: semua kandidat berputar mengelilingi titik terbaik"
    if k < 9:
        return "Dinamis: pusat spiral berpindah saat ditemukan solusi yang lebih baik"
    if k < 14:
        return "Jari-jari spiral menyusut — eksplorasi beralih menjadi eksploitasi"
    return "Kandidat semakin rapat di sekitar solusi optimal"


# ---------------- timeline frame ----------------
frames = []  # tiap item: dict state
for i in range(INTRO_FRAMES):
    frames.append(dict(phase="intro", t=i / (INTRO_FRAMES - 1)))
for i in range(HILITE_FRAMES):
    frames.append(dict(phase="hilite", t=i / (HILITE_FRAMES - 1)))
for k in range(N_ITER):
    for i in range(FRAMES_PER_ITER):
        frames.append(dict(phase="iter", k=k, t=smoothstep(i / (FRAMES_PER_ITER - 1))))
for i in range(END_FRAMES):
    frames.append(dict(phase="end", t=i / (END_FRAMES - 1)))

print(f"total frame: {len(frames)}  (~{len(frames)/FPS:.1f} detik)")

# ---------------- render ----------------
fig = plt.figure(figsize=(W, H), dpi=DPI)
fig.patch.set_facecolor(BG)

writer = imageio.get_writer(
    OUT, fps=FPS, codec="libx264", quality=8, pixelformat="yuv420p", macro_block_size=1
)

trail = []  # riwayat posisi utk jejak spiral

FULL_LIM = (-5.4, 5.4)
ZOOM_HALF = 1.35


def view_limits(fr):
    """zoom perlahan ke area konvergensi pada iterasi akhir"""
    if fr["phase"] == "iter":
        k = fr["k"]
        prog = (k + fr["t"]) / N_ITER
        z = smoothstep(np.clip((prog - 0.62) / 0.38, 0, 1))
    elif fr["phase"] == "end":
        z = 1.0
    else:
        z = 0.0
    half = (FULL_LIM[1] - FULL_LIM[0]) / 2 * (1 - z) + ZOOM_HALF * z
    cx = 0.0 * (1 - z) + OPT[0] * z
    cy = 0.0 * (1 - z) + OPT[1] * z
    return (cx - half, cx + half), (cy - half, cy + half)


for idx, fr in enumerate(frames):
    fig.clf()
    ax = fig.add_axes([0.045, 0.10, 0.62, 0.80])
    ax.set_facecolor(BG)

    # ---- state posisi ----
    phase = fr["phase"]
    if phase == "intro":
        pos = iters[0]
        center = None
        k_show, r_show = 0, r_at(0)
        n_vis = int(np.ceil(fr["t"] * N_POINTS))  # titik muncul satu per satu
        caption = "Inisialisasi: kandidat solusi tersebar acak di ruang pencarian"
    elif phase == "hilite":
        pos = iters[0]
        center = centers[0]
        k_show, r_show = 0, rvals[0]
        n_vis = N_POINTS
        caption = "Evaluasi: titik dengan nilai fungsi terbaik menjadi pusat spiral"
    elif phase == "iter":
        k = fr["k"]
        pos = interp_positions(k, fr["t"])
        center = centers[k]
        k_show, r_show = k + 1, rvals[k]
        n_vis = N_POINTS
        caption = caption_for(k)
    else:  # end
        pos = iters[-1]
        center = centers[-1]
        k_show, r_show = N_ITER, rvals[-1]
        n_vis = N_POINTS
        caption = "Konvergen! Seluruh kandidat menuju solusi optimal x*"

    if phase == "iter":
        trail.append(pos.copy())
        if len(trail) > TRAIL_LEN:
            trail.pop(0)

    xlim, ylim = view_limits(fr)

    # ---- kontur fungsi objektif ----
    ax.contourf(GX, GY, GZ, levels=LEVELS, cmap="viridis", alpha=0.28)
    ax.contour(GX, GY, GZ, levels=LEVELS, colors="white", linewidths=0.4, alpha=0.14)

    # ---- jejak spiral ----
    if len(trail) > 1:
        arr = np.array(trail)  # (T, N, 2)
        segs, alphas = [], []
        T = len(arr)
        for j in range(T - 1):
            a = (j + 1) / T
            for n in range(N_POINTS):
                segs.append([arr[j, n], arr[j + 1, n]])
                alphas.append(0.35 * a)
        lc = LineCollection(segs, colors=[(0.22, 0.78, 0.94, a) for a in alphas], linewidths=1.0)
        ax.add_collection(lc)

    # ---- titik kandidat ----
    ax.scatter(
        pos[:n_vis, 0], pos[:n_vis, 1], s=42, c=ACCENT,
        edgecolors="white", linewidths=0.6, zorder=5,
    )

    # ---- pusat spiral (titik terbaik) ----
    if center is not None:
        if phase == "hilite":
            pulse = 0.25 + 0.55 * (0.5 + 0.5 * np.sin(fr["t"] * 6 * np.pi))
            circ = plt.Circle(center, 0.55 + 0.25 * pulse, fill=False,
                              color=GOLD, lw=2.0, alpha=0.8 * (1 - fr["t"] * 0.3))
            ax.add_patch(circ)
        ax.scatter(*center, s=340, marker="*", c=GOLD,
                   edgecolors="white", linewidths=0.8, zorder=6)

    # ---- tanda solusi optimal di fase akhir ----
    if phase == "end":
        a = min(1.0, fr["t"] * 2.5)
        ax.scatter(*OPT, s=180, marker="+", c="white", linewidths=2, alpha=a, zorder=7)
        ax.annotate(
            "solusi optimal $x^*$", xy=OPT, xytext=(OPT[0] + 0.45, OPT[1] + 0.55),
            color="white", fontsize=13, alpha=a,
            arrowprops=dict(arrowstyle="->", color="white", alpha=a),
        )

    ax.set_xlim(*xlim)
    ax.set_ylim(*ylim)
    ax.set_aspect("equal")
    ax.tick_params(colors=DIM, labelsize=8)
    for s in ax.spines.values():
        s.set_color("#33404d")

    # ---- panel teks kanan ----
    fig.text(0.695, 0.86, "Dynamic Spiral\nOptimization", color=TXT,
             fontsize=21, fontweight="bold", va="top")
    fig.text(0.695, 0.70, "Aturan pembaruan:", color=DIM, fontsize=11)
    fig.text(0.695, 0.63,
             r"$x_i^{(k+1)} = x^{*} + r\,R(\theta)\,(x_i^{(k)} - x^{*})$",
             color=TXT, fontsize=13.5)
    fig.text(0.695, 0.555, r"$R(\theta)$ : rotasi   $\bullet$   $r<1$ : kontraksi",
             color=DIM, fontsize=10.5)

    fig.text(0.695, 0.46, f"Iterasi", color=DIM, fontsize=11)
    fig.text(0.87, 0.46, f"{k_show} / {N_ITER}", color=TXT, fontsize=12, fontweight="bold")
    fig.text(0.695, 0.41, f"Sudut rotasi θ", color=DIM, fontsize=11)
    fig.text(0.87, 0.41, "45°", color=TXT, fontsize=12, fontweight="bold")
    fig.text(0.695, 0.36, f"Kontraksi r (dinamis)", color=DIM, fontsize=11)
    fig.text(0.87, 0.36, f"{r_show:.3f} ↓", color=TXT, fontsize=12, fontweight="bold")
    fbest = feval(pos[:n_vis]).min() if n_vis > 0 else float("nan")
    fig.text(0.695, 0.31, f"f(x*) terbaik", color=DIM, fontsize=11)
    fig.text(0.87, 0.31, f"{fbest:.4f}", color=GOLD, fontsize=12, fontweight="bold")

    # legenda mini
    fig.text(0.695, 0.225, "●", color=ACCENT, fontsize=13)
    fig.text(0.715, 0.225, "kandidat solusi", color=TXT, fontsize=10.5)
    fig.text(0.695, 0.185, "★", color=GOLD, fontsize=13)
    fig.text(0.715, 0.185, "titik terbaik = pusat spiral", color=TXT, fontsize=10.5)

    # ---- bilah keterangan bawah ----
    fig.text(0.045, 0.032, caption, color=TXT, fontsize=13.5,
             bbox=dict(boxstyle="round,pad=0.45", fc="#16202b", ec="#33404d"))

    fig.canvas.draw()
    buf = np.asarray(fig.canvas.buffer_rgba())[:, :, :3]
    writer.append_data(buf)

    if idx % 100 == 0:
        print(f"frame {idx}/{len(frames)}")

writer.close()
plt.close(fig)
print("SELESAI:", OUT)
