File size: 1,921 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
"""
Denoising Diffusion Probabilistic Models — Training Objective (Simplified Loss)

Paper: https://arxiv.org/abs/2006.11239
Authors: Ho, Jain, Abbeel (2020)

Implements: L_simple from §3.4, Eq. 14

"We have shown that the variational bound... can be optimized with a
 simplified objective... which resembles denoising score matching..."

  L_simple = E_{t, x_0, ε} [ ||ε − ε_θ(√ᾱ_t x_0 + √(1−ᾱ_t) ε, t)||² ]

This is MSE between the true noise ε and the model's prediction ε_θ.
The expectation is over:
  - t ~ Uniform({1, ..., T})
  - x_0 ~ q(x_0)  (data distribution)
  - ε ~ N(0, I)

§3.4 — "We found it beneficial to sample quality (and simpler to implement)
to train on the following variant of the variational bound... L_simple"

§3.4 — "Algorithm 1" describes the training procedure that uses this loss.
"""

import torch
import torch.nn as nn


class DDPMLoss(nn.Module):
    """§3.4, Eq. 14 — Simplified training objective L_simple.

    Computes MSE between true noise and predicted noise:
        L = ||ε − ε_θ(x_t, t)||²

    This module handles only the loss computation. The caller (training loop)
    is responsible for sampling t, computing x_t from x_0, and calling the model.
    """

    def __init__(self):
        super().__init__()

    def forward(
        self,
        noise_pred: torch.Tensor,
        noise_true: torch.Tensor,
    ) -> torch.Tensor:
        """
        §3.4, Eq. 14 — L_simple = E[||ε − ε_θ(x_t, t)||²]

        Args:
            noise_pred: (batch, C, H, W) — predicted noise ε_θ(x_t, t)
            noise_true: (batch, C, H, W) — true noise ε ~ N(0, I)

        Returns:
            scalar — mean squared error loss
        """
        # §3.4 — Simple MSE between predicted and true noise
        # "equivalent to (a re-weighted variant of) the ELBO"
        return nn.functional.mse_loss(noise_pred, noise_true)