File size: 5,704 Bytes
f493668
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
"""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