"""
DDPM Animation: forward (добавление шума) + reverse (удаление шума)
@fminxyz Series 2, Post 1 — 4 марта 2026
1080x1080 px, 25 fps, ~20 сек
"""

import numpy as np
import matplotlib.pyplot as plt
import matplotlib.animation as animation
from matplotlib.patches import FancyArrowPatch

# ── Params ──────────────────────────────────────────────────────────────────
STEPS = 10          # diffusion steps shown
FPS = 2             # frames per second
HOLD = 2            # frames to hold at each step
IMG_SIZE = 64       # synthetic "image" resolution
OUTPUT = "Strategy/content/drafts/ddpm_animation.mp4"

rng = np.random.default_rng(42)

# ── Synthetic "image" (simple geometric pattern) ─────────────────────────────
def make_image(size=IMG_SIZE):
    img = np.zeros((size, size, 3))
    # gradient background
    for i in range(size):
        img[i, :, 0] = i / size * 0.8
        img[i, :, 2] = (size - i) / size * 0.8
    # white circle
    cx, cy = size // 2, size // 2
    r = size // 4
    y, x = np.ogrid[:size, :size]
    mask = (x - cx)**2 + (y - cy)**2 <= r**2
    img[mask] = [1.0, 0.9, 0.2]
    return np.clip(img, 0, 1)

x0 = make_image()

# ── DDPM forward: q(x_t | x_0) = N(√ᾱ_t * x_0, (1 - ᾱ_t) I) ────────────────
betas = np.linspace(0.02, 0.4, STEPS)
alphas = 1.0 - betas
alpha_bars = np.cumprod(alphas)

def forward(x0, t):
    ab = alpha_bars[t]
    eps = rng.standard_normal(x0.shape)
    return np.sqrt(ab) * x0 + np.sqrt(1 - ab) * eps

# Forward frames: x0 → x_T
forward_frames = [x0] + [forward(x0, t) for t in range(STEPS)]

# Reverse frames: x_T → x0  (approx: lerp back using true x0 as oracle)
def reverse_approx(xt, t, x0):
    ab = alpha_bars[t]
    frac = t / (STEPS - 1)
    return np.clip((1 - frac) * x0 + frac * xt, 0, 1)

reverse_frames = [forward_frames[-1]] + [
    reverse_approx(forward_frames[-1], STEPS - 1 - t, x0)
    for t in range(STEPS)
]
reverse_frames[-1] = x0  # ensure clean end

# ── Build frame sequence with holds ─────────────────────────────────────────
all_frames = []
labels = []

# Phase 1: Forward (→ noise)
for i, frame in enumerate(forward_frames):
    for _ in range(HOLD):
        all_frames.append(frame)
        if i == 0:
            labels.append("Оригинал")
        elif i == STEPS:
            labels.append("Чистый шум x_T")
        else:
            labels.append(f"Шаг t={i}: q(x_t|x_{{t-1}})")

# Phase 2: Reverse (→ clean)
for i, frame in enumerate(reverse_frames):
    for _ in range(HOLD):
        all_frames.append(frame)
        t_rev = STEPS - i
        if i == 0:
            labels.append("Чистый шум x_T")
        elif i == STEPS:
            labels.append("Восстановлено! ✓")
        else:
            labels.append(f"Обратный шаг t={t_rev}: p_θ(x_{{t-1}}|x_t)")

# ── Figure setup ─────────────────────────────────────────────────────────────
fig, ax = plt.subplots(figsize=(10.8, 10.8), dpi=100)
fig.patch.set_facecolor('#0d0d0d')
ax.set_facecolor('#0d0d0d')
ax.set_xticks([])
ax.set_yticks([])
for spine in ax.spines.values():
    spine.set_visible(False)

im = ax.imshow(all_frames[0], interpolation='bilinear', aspect='equal')

# Phase label (top)
phase_text = ax.text(
    0.5, 0.97, "FORWARD →  добавляем шум",
    transform=ax.transAxes, ha='center', va='top',
    fontsize=18, color='#ff6b6b', fontweight='bold',
    fontfamily='monospace'
)

# Step label (bottom)
step_text = ax.text(
    0.5, 0.03, labels[0],
    transform=ax.transAxes, ha='center', va='bottom',
    fontsize=13, color='#cccccc', fontfamily='monospace'
)

# Progress bar (in axes coordinates via transAxes)
from matplotlib.lines import Line2D
bar_bg = fig.add_artist(Line2D([0.05, 0.95], [0.04, 0.04],
                               transform=fig.transFigure,
                               color='#333333', linewidth=6, zorder=5))
progress_line = fig.add_artist(Line2D([0.05, 0.05], [0.04, 0.04],
                                      transform=fig.transFigure,
                                      color='#ff6b6b', linewidth=6, zorder=6))

n_total = len(all_frames)
n_fwd = len(forward_frames) * HOLD

def update(frame_idx):
    im.set_data(all_frames[frame_idx])
    step_text.set_text(labels[frame_idx])

    # Progress bar
    frac = frame_idx / (n_total - 1)
    progress_line.set_xdata([0.05, 0.05 + frac * 0.90])

    # Phase
    if frame_idx < n_fwd:
        phase_text.set_text("FORWARD  →  добавляем шум")
        phase_text.set_color('#ff6b6b')
        progress_line.set_color('#ff6b6b')
    else:
        phase_text.set_text("REVERSE  ←  убираем шум")
        phase_text.set_color('#6bffb8')
        progress_line.set_color('#6bffb8')

    return im, step_text, phase_text, progress_line

ani = animation.FuncAnimation(
    fig, update, frames=n_total,
    interval=1000 // FPS, blit=True
)

writer = animation.FFMpegWriter(fps=FPS, bitrate=1200,
                                extra_args=['-vcodec', 'libx264', '-pix_fmt', 'yuv420p'])
ani.save(OUTPUT, writer=writer)
print(f"Saved: {OUTPUT}  ({n_total} frames @ {FPS} fps)")
plt.close()
