File size: 8,171 Bytes
256c9c2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
"""
Denoising Diffusion Probabilistic Models — Shared Utilities

Paper: https://arxiv.org/abs/2006.11239
Implements: Noise schedule (§4), forward process q(x_t|x_0) (§3, Eq. 4),
            and reverse sampling (Algorithm 2).

Section references:
  §2 — Forward and reverse diffusion processes
  §3 — Training objective derivation
  §4 — Noise schedule specification
  Algorithm 1 — Training
  Algorithm 2 — Sampling
"""

import math
from typing import Dict, Optional

import torch
import torch.nn as nn


def linear_noise_schedule(
    timesteps: int = 1000,
    beta_start: float = 1e-4,
    beta_end: float = 0.02,
) -> Dict[str, torch.Tensor]:
    """§4 — Linear variance schedule.

    "We set the forward process variances to constants increasing linearly
     from β_1 = 10^-4 to β_T = 0.02."

    Precomputes all quantities needed for training and sampling.

    Args:
        timesteps: §4 — T = 1000
        beta_start: §4 — β_1 = 10^-4
        beta_end: §4 — β_T = 0.02

    Returns:
        Dict with precomputed schedule tensors
    """
    # §4 — linear schedule
    betas = torch.linspace(beta_start, beta_end, timesteps)

    # §2 — α_t = 1 - β_t
    alphas = 1.0 - betas

    # §2 — α̅_t = Π_{s=1}^{t} α_s  (cumulative product)
    alphas_cumprod = torch.cumprod(alphas, dim=0)

    # α̅_{t-1} — needed for posterior q(x_{t-1} | x_t, x_0)
    alphas_cumprod_prev = torch.cat([torch.tensor([1.0]), alphas_cumprod[:-1]])

    # Precompute quantities for q(x_t | x_0) = N(x_t; √α̅_t x_0, (1-α̅_t)I)
    sqrt_alphas_cumprod = torch.sqrt(alphas_cumprod)
    sqrt_one_minus_alphas_cumprod = torch.sqrt(1.0 - alphas_cumprod)

    # Precompute for sampling (Algorithm 2)
    sqrt_recip_alphas = 1.0 / torch.sqrt(alphas)

    # §3.2, Eq. 7 — Posterior variance β̃_t = β_t * (1 - α̅_{t-1}) / (1 - α̅_t)
    posterior_variance = betas * (1.0 - alphas_cumprod_prev) / (1.0 - alphas_cumprod)

    return {
        "betas": betas,
        "alphas": alphas,
        "alphas_cumprod": alphas_cumprod,
        "alphas_cumprod_prev": alphas_cumprod_prev,
        "sqrt_alphas_cumprod": sqrt_alphas_cumprod,
        "sqrt_one_minus_alphas_cumprod": sqrt_one_minus_alphas_cumprod,
        "sqrt_recip_alphas": sqrt_recip_alphas,
        "posterior_variance": posterior_variance,
    }


def q_sample(
    x_0: torch.Tensor,
    t: torch.Tensor,
    schedule: Dict[str, torch.Tensor],
    noise: Optional[torch.Tensor] = None,
) -> torch.Tensor:
    """§3, Eq. 4 — Forward process: sample x_t from q(x_t | x_0).

    "A notable property is that we can sample x_t at any arbitrary time step t
     in closed form: q(x_t | x_0) = N(x_t; √α̅_t x_0, (1 - α̅_t)I)"

    x_t = √α̅_t * x_0 + √(1 - α̅_t) * ε, where ε ~ N(0, I)

    Args:
        x_0: (batch, C, H, W) — clean images
        t: (batch,) — timestep indices
        schedule: precomputed noise schedule
        noise: optional pre-sampled noise (for reproducibility)

    Returns:
        x_t: (batch, C, H, W) — noisy images at timestep t
    """
    if noise is None:
        noise = torch.randn_like(x_0)

    # Extract schedule values for timestep t, reshape for broadcasting
    sqrt_alpha_cumprod = schedule["sqrt_alphas_cumprod"][t]          # (batch,)
    sqrt_one_minus_alpha_cumprod = schedule["sqrt_one_minus_alphas_cumprod"][t]  # (batch,)

    # Reshape for broadcasting with (batch, C, H, W)
    sqrt_alpha_cumprod = sqrt_alpha_cumprod.view(-1, 1, 1, 1)
    sqrt_one_minus_alpha_cumprod = sqrt_one_minus_alpha_cumprod.view(-1, 1, 1, 1)

    # §3, Eq. 4 — x_t = √α̅_t * x_0 + √(1 - α̅_t) * ε
    return sqrt_alpha_cumprod * x_0 + sqrt_one_minus_alpha_cumprod * noise


@torch.no_grad()
def p_sample(
    model: nn.Module,
    x_t: torch.Tensor,
    t: torch.Tensor,
    t_index: int,
    schedule: Dict[str, torch.Tensor],
) -> torch.Tensor:
    """Algorithm 2, lines 3-4 — Single reverse step: sample x_{t-1} from p_θ(x_{t-1} | x_t).

    "x_{t-1} = 1/√α_t * (x_t - β_t/√(1-α̅_t) * ε_θ(x_t, t)) + σ_t * z"

    where z ~ N(0, I) if t > 1, else z = 0.

    Args:
        model: noise prediction network ε_θ
        x_t: (batch, C, H, W) — current noisy sample
        t: (batch,) — current timestep (as tensor for model input)
        t_index: integer timestep (for indexing schedule)
        schedule: precomputed noise schedule

    Returns:
        x_{t-1}: (batch, C, H, W) — denoised sample one step
    """
    # Predict noise ε_θ(x_t, t)
    predicted_noise = model(x_t, t)  # (batch, C, H, W)

    # Extract schedule values
    beta_t = schedule["betas"][t_index]
    sqrt_recip_alpha_t = schedule["sqrt_recip_alphas"][t_index]
    sqrt_one_minus_alpha_cumprod_t = schedule["sqrt_one_minus_alphas_cumprod"][t_index]

    # Algorithm 2, line 4 — Compute mean of p_θ(x_{t-1} | x_t)
    # μ_θ = 1/√α_t * (x_t - β_t/√(1-α̅_t) * ε_θ(x_t, t))
    mean = sqrt_recip_alpha_t * (
        x_t - beta_t / sqrt_one_minus_alpha_cumprod_t * predicted_noise
    )

    if t_index == 0:
        # Algorithm 2, line 3 — z = 0 when t = 1 (final step)
        return mean
    else:
        # Algorithm 2, line 3 — z ~ N(0, I) when t > 1
        # §3.4 — σ²_t = β_t (fixed small variance)
        sigma_t = torch.sqrt(schedule["betas"][t_index])
        noise = torch.randn_like(x_t)
        return mean + sigma_t * noise


@torch.no_grad()
def sample(
    model: nn.Module,
    schedule: Dict[str, torch.Tensor],
    image_shape: tuple,
    device: torch.device,
) -> torch.Tensor:
    """Algorithm 2 — Full reverse sampling process.

    "Algorithm 2 Sampling
     1: x_T ~ N(0, I)
     2: for t = T, ..., 1 do
     3:   z ~ N(0, I) if t > 1, else z = 0
     4:   x_{t-1} = 1/√α_t (x_t - β_t/√(1-α̅_t) ε_θ(x_t, t)) + σ_t z
     5: end for
     6: return x_0"

    Args:
        model: noise prediction network ε_θ (should be in eval mode, ideally EMA weights)
        schedule: precomputed noise schedule
        image_shape: (batch, C, H, W) — shape of images to generate
        device: torch device

    Returns:
        x_0: (batch, C, H, W) — generated images
    """
    model.eval()
    timesteps = len(schedule["betas"])

    # Algorithm 2, line 1: x_T ~ N(0, I)
    x = torch.randn(image_shape, device=device)

    # Algorithm 2, lines 2-5: reverse iterate from t=T to t=1
    for t_index in reversed(range(timesteps)):
        t = torch.full((image_shape[0],), t_index, device=device, dtype=torch.long)
        x = p_sample(model, x, t, t_index, schedule)

    return x


class EMA:
    """§4 — Exponential Moving Average of model parameters.

    "We report sample quality metrics using an exponential moving average (EMA)
     of model parameters with a decay factor of 0.9999."

    The EMA weights are used for sampling/evaluation, not for training.
    """

    def __init__(self, model: nn.Module, decay: float = 0.9999):
        """
        Args:
            model: the model to track
            decay: §4 — "decay factor of 0.9999"
        """
        self.decay = decay
        self.shadow = {}
        for name, param in model.named_parameters():
            if param.requires_grad:
                self.shadow[name] = param.data.clone()

    @torch.no_grad()
    def update(self, model: nn.Module):
        """Update EMA weights after each training step."""
        for name, param in model.named_parameters():
            if param.requires_grad and name in self.shadow:
                self.shadow[name].mul_(self.decay).add_(
                    param.data, alpha=1.0 - self.decay
                )

    def apply(self, model: nn.Module):
        """Load EMA weights into model (for evaluation/sampling)."""
        for name, param in model.named_parameters():
            if name in self.shadow:
                param.data.copy_(self.shadow[name])

    def restore(self, model: nn.Module):
        """Restore original model weights (after evaluation)."""
        # NOTE: This requires storing original weights separately.
        # The caller should save model.state_dict() before calling apply().
        pass