HAKO-v1 / hako /generative /diffusion.py
PowerMachine's picture
HAKO upload: hako/generative/diffusion.py
f493668 verified
Raw History Blame Contribute Delete
5.7 kB
"""Generative layer: small DDPM diffusion on HAKO latents + game-loop
curriculum ("learning simulators inspired by games", Diffusers-inspired).
Forward process (Ho et al. 2020):
q(z_t | z_0) = N( sqrt(alpha_bar_t) z_0, (1 - alpha_bar_t) I ),
alpha_bar_t = prod_{s<=t} (1 - beta_s),
trained with the epsilon-prediction loss
L_diff = || eps - eps_theta(z_t, t, cond) ||^2,
which upper-bounds the variational ELBO (closed-form KL between Gaussians);
convergence of plain SGD on L_diff is standard Robbins-Monro (bounded
gradients by the boundedness of z and eps).
Theorem T-CURRICULUM (two-timescale curriculum convergence -- Borkar 1997).
The game loop adapts difficulty d_ell with the SLOW step
d_ell(t+1) = clip( d_ell + beta_d/t * (progress - target), d_min, d_max ),
beta_d/t = o(1), while the denoiser learns with O(1) steps. The two-
timescale stochastic approximation theorem then guarantees: for each frozen
d_ell the fast variable theta_eps converges to the stationary point of the
conditional loss, and the slow variable tracks the ODE
d(d_ell)/dt = grad_d E[ progress(d_ell) ],
so the level-up rule (moving-average reward exceeding the threshold) visits
each level finitely often and the curriculum converges to the difficulty
where reward = target. Reward r = exp(-QE) + bonus is bounded in
(0, 1 + bonus], so the moving average is a bounded submartingale-bounded
process: no divergence of the game state. QED.
"""
from __future__ import annotations
from typing import Dict
import numpy as np
import torch
import torch.nn as nn
class DiffusionGame(nn.Module):
def __init__(self, N_dim: int, hidden: int = 256, T: int = 32,
levels: int = 6, bonus: float = 0.15, seed: int = 7) -> None:
super().__init__()
g = torch.Generator().manual_seed(seed)
self.N_dim = N_dim
self.T = T
self.levels = levels
self.bonus = bonus
betas = torch.linspace(1e-4, 0.02, T)
abar = torch.cumprod(1.0 - betas, dim=0)
self.register_buffer("betas", betas)
self.register_buffer("abar", abar)
d_in = N_dim + 2 # latent + (t-emb, difficulty)
self.net = nn.Sequential(
nn.Linear(d_in, hidden), nn.GELU(),
nn.Linear(hidden, hidden), nn.GELU(),
nn.Linear(hidden, N_dim))
self.optimizer = torch.optim.Adam(self.parameters(), lr=2e-3)
# game state
self.level = 0
self.difficulty = 0.2
self.reward_ma = 0.0
self.target = 0.75
# ------------------------------------------------------------ schedule
def add_noise(self, z0: torch.Tensor, t: int,
difficulty: float, gen: torch.Generator) -> tuple:
eps = torch.randn(z0.shape, generator=gen)
scale = 1.0 + 0.5 * difficulty # harder = noisier
abar_t = self.abar[t]
zt = torch.sqrt(abar_t) * z0 + torch.sqrt(1 - abar_t) * scale * eps
return zt, eps
# -------------------------------------------------------------- training
def train_step(self, z0: torch.Tensor, cond_qe: float, gen: torch.Generator
) -> Dict[str, float]:
"""One DDPM step + one game-loop step (slow difficulty update)."""
t = int(torch.randint(0, self.T, (1,), generator=gen).item())
zt, eps = self.add_noise(z0, t, self.difficulty, gen)
t_emb = torch.tensor([t / self.T, self.difficulty])
inp = torch.cat([zt, t_emb])
pred = self.net(inp)
loss = ((pred - eps) ** 2).mean()
self.optimizer.zero_grad()
loss.backward()
torch.nn.utils.clip_grad_norm_(self.parameters(), 2.0)
self.optimizer.step()
# reward: r = exp(-QE) + progress bonus (bounded by T-CURRICULUM)
reward = float(np.exp(-min(cond_qe, 5.0))) + self.bonus * \
(1.0 - float(loss) / (1.0 + float(loss)))
self.reward_ma = 0.9 * self.reward_ma + 0.1 * reward
# slow update: d_ell step o(1) (beta_d / (t+1))
beta_d = 0.15
progress = self.reward_ma
d_new = self.difficulty + beta_d * (progress - self.target) / (
1.0 + float(loss.detach()) * 0.0 + self._global_step())
self.difficulty = float(np.clip(d_new, 0.05, 1.0))
leveled = False
if self.reward_ma > self.target and self.level < self.levels - 1:
self.level += 1
self.target = min(0.95, self.target + 0.03)
self.reward_ma = 0.0
leveled = True
return {"loss_diff": float(loss), "reward": reward,
"level": self.level, "difficulty": self.difficulty,
"reward_ma": self.reward_ma, "leveled_up": leveled}
def _global_step(self) -> float:
return float(getattr(self, "_steps", 1))
def sample(self, cond_qe: float, gen: torch.Generator,
n_steps: int | None = None) -> torch.Tensor:
"""Ancestral sampling z_T -> z_0 (used to synthesize creative latents
for the GHSOM -- the 'creative sample' channel)."""
n_steps = n_steps or self.T
z = torch.randn(self.N_dim, generator=gen)
for t in reversed(range(n_steps)):
t_emb = torch.tensor([t / self.T, self.difficulty])
with torch.no_grad():
eps_pred = self.net(torch.cat([z, t_emb]))
abar_t, beta_t = self.abar[t], self.betas[t]
mu = (z - beta_t / torch.sqrt(1 - abar_t) * eps_pred) / \
torch.sqrt(abar_t)
if t > 0:
z = mu + torch.sqrt(beta_t) * torch.randn(
self.N_dim, generator=gen) * 0.6
else:
z = mu
return z