# Source: content/notes/deep-learning/generative-models.md
# Independent CPU example; use the curriculum environment.
# See /notes/ml/#example-environment or /notes/deep-learning/#example-environment.

import torch
from torch import nn

torch.manual_seed(32)
torch.set_num_threads(1)
centers = torch.tensor([[-1., -1.], [-1., 1.], [1., -1.], [1., 1.]])
data = centers[torch.randint(4, (512,))] + 0.1 * torch.randn(512, 2)
steps = 40
beta = torch.linspace(0.01, 0.20, steps)
alpha = 1 - beta
abar = alpha.cumprod(0)
previous_abar = torch.cat((torch.ones(1), abar[:-1]))
posterior_variance = beta * (1 - previous_abar) / (1 - abar)
net = nn.Sequential(nn.Linear(3, 64), nn.SiLU(), nn.Linear(64, 64),
                    nn.SiLU(), nn.Linear(64, 2))
optimizer = torch.optim.Adam(net.parameters(), lr=0.005)

def predict(noisy, times):
    return net(torch.cat((noisy, times[:, None].float() / (steps - 1)), dim=1))

def denoising_loss(clean, times, noise):
    a = abar[times, None]
    noisy = a.sqrt() * clean + (1 - a).sqrt() * noise
    return nn.functional.mse_loss(predict(noisy, times), noise)

eval_times = torch.randint(steps, (len(data),))
eval_noise = torch.randn_like(data)
initial = denoising_loss(data, eval_times, eval_noise).item()
for _ in range(250):
    ids = torch.randint(len(data), (128,))
    times = torch.randint(steps, (128,))
    optimizer.zero_grad(set_to_none=True)
    loss = denoising_loss(data[ids], times, torch.randn(128, 2))
    loss.backward()
    optimizer.step()
net.eval()
with torch.no_grad():
    final = denoising_loss(data, eval_times, eval_noise).item()
    sample = torch.randn(256, 2)
    for t in reversed(range(steps)):
        times = torch.full((len(sample),), t, dtype=torch.long)
        epsilon = predict(sample, times)
        mean = (sample - beta[t] * epsilon / (1 - abar[t]).sqrt()) / alpha[t].sqrt()
        sample = mean if t == 0 else mean + posterior_variance[t].sqrt() * torch.randn_like(sample)
    counts = torch.bincount(torch.cdist(sample, centers).argmin(1), minlength=4)
assert posterior_variance[0] == 0
assert final < initial * 0.85
assert sample.shape == (256, 2) and torch.isfinite(sample).all()
print({"initial_noise_mse": initial, "final_noise_mse": final,
       "terminal_signal_fraction": abar[-1].item(), "mode_counts": counts.tolist()})
