File size: 7,081 Bytes
59630ba | 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 | """
A very minimal implementation of continuous-time diffusion models. For compatibility with other modules,
sampling schedules are still implemented in discrete time.
"""
from abc import ABC, abstractmethod
from typing import Optional
from omegaconf import DictConfig
import torch
from torch import nn
from torch.nn import functional as F
from .discrete_diffusion import DiscreteDiffusion, ModelPrediction
class ContinuousNoiseSchedule(nn.Module, ABC):
"""
An abstract class for continuous noise schedule that is compatible with continuous-time diffusion models.
"""
@classmethod
def from_config(cls, cfg: DictConfig):
match cfg.name:
case "cosine":
return CosineNoiseSchedule(cfg)
case _:
raise ValueError(f"unknown noise schedule {cfg.name}")
@abstractmethod
def forward(self, t: torch.Tensor) -> torch.Tensor:
"""Given the timestep t within [0, 1], return the logSNR value at that timestep."""
raise NotImplementedError
@property
@abstractmethod
def max_logsnr(self) -> torch.Tensor:
"""Return the maximum logSNR value."""
raise NotImplementedError
@property
@abstractmethod
def min_logsnr(self) -> torch.Tensor:
"""Return the minimum logSNR value."""
raise NotImplementedError
class CosineNoiseSchedule(ContinuousNoiseSchedule):
"""
Cosine noise schedule that can be shifted from base resolution to target resolution,
proposed in Simple Diffusion (2023, https://arxiv.org/abs/2301.11093).
Here, `shift` should be set to `base_resolution / target_resolution`.
"""
def __init__(self, cfg: DictConfig):
super().__init__()
logsnr_min, logsnr_max = cfg.get("logsnr_min", -15.0), cfg.get(
"logsnr_max", 15.0
)
shift = cfg.get("shift", 1.0)
self.register_buffer(
"t_min",
torch.atan(torch.exp(-0.5 * torch.tensor(logsnr_max, dtype=torch.float32))),
persistent=False,
)
self.register_buffer(
"t_max",
torch.atan(torch.exp(-0.5 * torch.tensor(logsnr_min, dtype=torch.float32))),
persistent=False,
)
self.register_buffer(
"shift",
2 * torch.log(torch.tensor(shift, dtype=torch.float32)),
persistent=False,
)
def forward(self, t: torch.Tensor) -> torch.Tensor:
return (
-2 * torch.log(torch.tan(self.t_min + t * (self.t_max - self.t_min)))
+ self.shift
)
@property
def max_logsnr(self) -> torch.Tensor:
return self.forward(
torch.tensor(0.0, dtype=torch.float32, device=self.shift.device)
)
@property
def min_logsnr(self) -> torch.Tensor:
return self.forward(
torch.tensor(1.0, dtype=torch.float32, device=self.shift.device)
)
class ContinuousDiffusion(DiscreteDiffusion):
def __init__(
self,
cfg: DictConfig,
backbone_cfg: DictConfig,
x_shape: torch.Size,
max_tokens: int,
external_cond_dim: int,
):
super().__init__(cfg, backbone_cfg, x_shape, max_tokens, external_cond_dim)
assert (
self.objective == "pred_v" and self.loss_weighting.strategy == "sigmoid"
), "ContinuousDiffusion only supports 'pred_v' objective and 'sigmoid' loss weighting"
self.precond_scale = cfg.precond_scale
self.sigmoid_bias = cfg.loss_weighting.sigmoid_bias
def _build_buffer(self):
super()._build_buffer()
self.training_schedule = ContinuousNoiseSchedule.from_config(
self.cfg.training_schedule
)
def model_predictions(self, x, k, external_cond=None, external_cond_mask=None,**kwargs):
model_output = self.model(
x, self.precond_scale * self.logsnr[k], external_cond, external_cond_mask,**kwargs
)
if self.objective == "pred_noise":
pred_noise = torch.clamp(model_output, -self.clip_noise, self.clip_noise)
x_start = self.predict_start_from_noise(x, k, pred_noise)
elif self.objective == "pred_x0":
x_start = model_output
pred_noise = self.predict_noise_from_start(x, k, x_start)
elif self.objective == "pred_v":
v = model_output
x_start = self.predict_start_from_v(x, k, v)
pred_noise = self.predict_noise_from_v(x, k, v)
model_pred = ModelPrediction(pred_noise, x_start, model_output)
return model_pred
def forward(
self,
x: torch.Tensor,
external_cond: Optional[torch.Tensor],
k: torch.Tensor,
**kwargs
):
logsnr = self.training_schedule(k)
noise = torch.randn_like(x)
noise = torch.clamp(noise, -self.clip_noise, self.clip_noise)
alpha_t = self.add_shape_channels(torch.sigmoid(logsnr).sqrt())
sigma_t = self.add_shape_channels(torch.sigmoid(-logsnr).sqrt())
x_t = alpha_t * x + sigma_t * noise
# v-prediction
v_pred = self.model(x_t, self.precond_scale * logsnr, external_cond,**kwargs)
noise_pred = alpha_t * v_pred + sigma_t * x_t
x_pred = alpha_t * x_t - sigma_t * v_pred
loss = F.mse_loss(noise_pred, noise.detach(), reduction="none")
# sigmoid loss weighting
# proposed by Kingma & Gao (2023, https://arxiv.org/abs/2303.00848)
# further studied in Simple Diffusion 2 (2024, https://arxiv.org/abs/2410.19324)
loss_weight = torch.sigmoid(self.sigmoid_bias - logsnr)
loss_weight = self.add_shape_channels(loss_weight)
loss = loss * loss_weight
return x_pred, loss
def forward_with_alignment_loss(
self,
x: torch.Tensor,
external_cond: Optional[torch.Tensor],
k: torch.Tensor,
):
logsnr = self.training_schedule(k)
noise = torch.randn_like(x)
noise = torch.clamp(noise, -self.clip_noise, self.clip_noise)
alpha_t = self.add_shape_channels(torch.sigmoid(logsnr).sqrt())
sigma_t = self.add_shape_channels(torch.sigmoid(-logsnr).sqrt())
x_t = alpha_t * x + sigma_t * noise
# v-prediction
# import pdb; pdb.set_trace()
## START modification for alignment loss
v_pred,latents_list = self.model(x_t, self.precond_scale * logsnr, external_cond,return_latents=True)
noise_pred = alpha_t * v_pred + sigma_t * x_t
x_pred = alpha_t * x_t - sigma_t * v_pred
loss = F.mse_loss(noise_pred, noise.detach(), reduction="none")
# sigmoid loss weighting
# proposed by Kingma & Gao (2023, https://arxiv.org/abs/2303.00848)
# further studied in Simple Diffusion 2 (2024, https://arxiv.org/abs/2410.19324)
loss_weight = torch.sigmoid(self.sigmoid_bias - logsnr)
loss_weight = self.add_shape_channels(loss_weight)
loss = loss * loss_weight
return x_pred, loss , latents_list |