Skip to content

Commit ba21735

Browse files
committed
DDPM training example
1 parent 2d1f7de commit ba21735

9 files changed

Lines changed: 183 additions & 50 deletions

File tree

src/diffusers/__init__.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -9,6 +9,6 @@
99
from .models.unet_glide import GLIDESuperResUNetModel, GLIDETextToImageUNetModel
1010
from .models.unet_ldm import UNetLDMModel
1111
from .pipeline_utils import DiffusionPipeline
12-
from .pipelines import DDIM, DDPM, GLIDE, LatentDiffusion, BDDMPipeline
12+
from .pipelines import DDIM, DDPM, GLIDE, BDDMPipeline, LatentDiffusion
1313
from .schedulers import DDIMScheduler, DDPMScheduler, SchedulerMixin
1414
from .schedulers.classifier_free_guidance import ClassifierFreeGuidanceScheduler

src/diffusers/configuration_utils.py

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -225,11 +225,11 @@ def _dict_from_json_file(cls, json_file: Union[str, os.PathLike]):
225225
text = reader.read()
226226
return json.loads(text)
227227

228-
def __eq__(self, other):
229-
return self.__dict__ == other.__dict__
228+
# def __eq__(self, other):
229+
# return self.__dict__ == other.__dict__
230230

231-
def __repr__(self):
232-
return f"{self.__class__.__name__} {self.to_json_string()}"
231+
# def __repr__(self):
232+
# return f"{self.__class__.__name__} {self.to_json_string()}"
233233

234234
@property
235235
def config(self) -> Dict[str, Any]:
Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,5 @@
1+
from .pipeline_bddm import BDDMPipeline
12
from .pipeline_ddim import DDIM
23
from .pipeline_ddpm import DDPM
34
from .pipeline_glide import GLIDE
45
from .pipeline_latent_diffusion import LatentDiffusion
5-
from .pipeline_bddm import BDDMPipeline

src/diffusers/pipelines/conversion_glide.py

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -97,7 +97,9 @@
9797

9898
superres_model.load_state_dict(ups_state_dict, strict=False)
9999

100-
upscale_scheduler = DDIMScheduler(timesteps=1000, beta_schedule="linear", beta_start=0.0001, beta_end=0.02, tensor_format="pt")
100+
upscale_scheduler = DDIMScheduler(
101+
timesteps=1000, beta_schedule="linear", beta_start=0.0001, beta_end=0.02, tensor_format="pt"
102+
)
101103

102104
glide = GLIDE(
103105
text_unet=text2im_model,

src/diffusers/pipelines/pipeline_bddm.py

Lines changed: 44 additions & 35 deletions
Original file line numberDiff line numberDiff line change
@@ -13,14 +13,16 @@
1313

1414

1515
import math
16+
1617
import numpy as np
1718
import torch
1819
import torch.nn as nn
1920
import torch.nn.functional as F
21+
2022
import tqdm
2123

22-
from ..modeling_utils import ModelMixin
2324
from ..configuration_utils import ConfigMixin
25+
from ..modeling_utils import ModelMixin
2426
from ..pipeline_utils import DiffusionPipeline
2527

2628

@@ -46,8 +48,7 @@ def calc_diffusion_step_embedding(diffusion_steps, diffusion_step_embed_dim_in):
4648
_embed = np.log(10000) / (half_dim - 1)
4749
_embed = torch.exp(torch.arange(half_dim) * -_embed).cuda()
4850
_embed = diffusion_steps * _embed
49-
diffusion_step_embed = torch.cat((torch.sin(_embed),
50-
torch.cos(_embed)), 1)
51+
diffusion_step_embed = torch.cat((torch.sin(_embed), torch.cos(_embed)), 1)
5152
return diffusion_step_embed
5253

5354

@@ -67,8 +68,7 @@ class Conv(nn.Module):
6768
def __init__(self, in_channels, out_channels, kernel_size=3, dilation=1):
6869
super().__init__()
6970
self.padding = dilation * (kernel_size - 1) // 2
70-
self.conv = nn.Conv1d(in_channels, out_channels, kernel_size,
71-
dilation=dilation, padding=self.padding)
71+
self.conv = nn.Conv1d(in_channels, out_channels, kernel_size, dilation=dilation, padding=self.padding)
7272
self.conv = nn.utils.weight_norm(self.conv)
7373
nn.init.kaiming_normal_(self.conv.weight)
7474

@@ -94,24 +94,20 @@ def forward(self, x):
9494
# every residual block (named residual layer in paper)
9595
# contains one noncausal dilated conv
9696
class ResidualBlock(nn.Module):
97-
def __init__(self, res_channels, skip_channels, dilation,
98-
diffusion_step_embed_dim_out):
97+
def __init__(self, res_channels, skip_channels, dilation, diffusion_step_embed_dim_out):
9998
super().__init__()
10099
self.res_channels = res_channels
101100

102101
# Use a FC layer for diffusion step embedding
103102
self.fc_t = nn.Linear(diffusion_step_embed_dim_out, self.res_channels)
104103

105104
# Dilated conv layer
106-
self.dilated_conv_layer = Conv(self.res_channels, 2 * self.res_channels,
107-
kernel_size=3, dilation=dilation)
105+
self.dilated_conv_layer = Conv(self.res_channels, 2 * self.res_channels, kernel_size=3, dilation=dilation)
108106

109107
# Add mel spectrogram upsampler and conditioner conv1x1 layer
110108
self.upsample_conv2d = nn.ModuleList()
111109
for s in [16, 16]:
112-
conv_trans2d = nn.ConvTranspose2d(1, 1, (3, 2 * s),
113-
padding=(1, s // 2),
114-
stride=(1, s))
110+
conv_trans2d = nn.ConvTranspose2d(1, 1, (3, 2 * s), padding=(1, s // 2), stride=(1, s))
115111
conv_trans2d = nn.utils.weight_norm(conv_trans2d)
116112
nn.init.kaiming_normal_(conv_trans2d.weight)
117113
self.upsample_conv2d.append(conv_trans2d)
@@ -157,7 +153,7 @@ def forward(self, input_data):
157153
h += mel_spec
158154

159155
# Gated-tanh nonlinearity
160-
out = torch.tanh(h[:, :self.res_channels, :]) * torch.sigmoid(h[:, self.res_channels:, :])
156+
out = torch.tanh(h[:, : self.res_channels, :]) * torch.sigmoid(h[:, self.res_channels :, :])
161157

162158
# Residual and skip outputs
163159
res = self.res_conv(out)
@@ -169,10 +165,16 @@ def forward(self, input_data):
169165

170166

171167
class ResidualGroup(nn.Module):
172-
def __init__(self, res_channels, skip_channels, num_res_layers, dilation_cycle,
173-
diffusion_step_embed_dim_in,
174-
diffusion_step_embed_dim_mid,
175-
diffusion_step_embed_dim_out):
168+
def __init__(
169+
self,
170+
res_channels,
171+
skip_channels,
172+
num_res_layers,
173+
dilation_cycle,
174+
diffusion_step_embed_dim_in,
175+
diffusion_step_embed_dim_mid,
176+
diffusion_step_embed_dim_out,
177+
):
176178
super().__init__()
177179
self.num_res_layers = num_res_layers
178180
self.diffusion_step_embed_dim_in = diffusion_step_embed_dim_in
@@ -185,16 +187,19 @@ def __init__(self, res_channels, skip_channels, num_res_layers, dilation_cycle,
185187
self.residual_blocks = nn.ModuleList()
186188
for n in range(self.num_res_layers):
187189
self.residual_blocks.append(
188-
ResidualBlock(res_channels, skip_channels,
189-
dilation=2 ** (n % dilation_cycle),
190-
diffusion_step_embed_dim_out=diffusion_step_embed_dim_out))
190+
ResidualBlock(
191+
res_channels,
192+
skip_channels,
193+
dilation=2 ** (n % dilation_cycle),
194+
diffusion_step_embed_dim_out=diffusion_step_embed_dim_out,
195+
)
196+
)
191197

192198
def forward(self, input_data):
193199
x, mel_spectrogram, diffusion_steps = input_data
194200

195201
# Embed diffusion step t
196-
diffusion_step_embed = calc_diffusion_step_embedding(
197-
diffusion_steps, self.diffusion_step_embed_dim_in)
202+
diffusion_step_embed = calc_diffusion_step_embedding(diffusion_steps, self.diffusion_step_embed_dim_in)
198203
diffusion_step_embed = swish(self.fc_t1(diffusion_step_embed))
199204
diffusion_step_embed = swish(self.fc_t2(diffusion_step_embed))
200205

@@ -239,20 +244,24 @@ def __init__(
239244
diffusion_step_embed_dim_out=diffusion_step_embed_dim_out,
240245
)
241246

242-
243247
# Initial conv1x1 with relu
244248
self.init_conv = nn.Sequential(Conv(in_channels, res_channels, kernel_size=1), nn.ReLU(inplace=False))
245249
# All residual layers
246-
self.residual_layer = ResidualGroup(res_channels,
247-
skip_channels,
248-
num_res_layers,
249-
dilation_cycle,
250-
diffusion_step_embed_dim_in,
251-
diffusion_step_embed_dim_mid,
252-
diffusion_step_embed_dim_out)
250+
self.residual_layer = ResidualGroup(
251+
res_channels,
252+
skip_channels,
253+
num_res_layers,
254+
dilation_cycle,
255+
diffusion_step_embed_dim_in,
256+
diffusion_step_embed_dim_mid,
257+
diffusion_step_embed_dim_out,
258+
)
253259
# Final conv1x1 -> relu -> zeroconv1x1
254-
self.final_conv = nn.Sequential(Conv(skip_channels, skip_channels, kernel_size=1),
255-
nn.ReLU(inplace=False), ZeroConv1d(skip_channels, out_channels))
260+
self.final_conv = nn.Sequential(
261+
Conv(skip_channels, skip_channels, kernel_size=1),
262+
nn.ReLU(inplace=False),
263+
ZeroConv1d(skip_channels, out_channels),
264+
)
256265

257266
def forward(self, input_data):
258267
audio, mel_spectrogram, diffusion_steps = input_data
@@ -267,12 +276,12 @@ def __init__(self, diffwave, noise_scheduler):
267276
super().__init__()
268277
noise_scheduler = noise_scheduler.set_format("pt")
269278
self.register_modules(diffwave=diffwave, noise_scheduler=noise_scheduler)
270-
279+
271280
@torch.no_grad()
272281
def __call__(self, mel_spectrogram, generator):
273282
if torch_device is None:
274283
torch_device = "cuda" if torch.cuda.is_available() else "cpu"
275-
284+
276285
self.diffwave.to(torch_device)
277286

278287
audio_length = mel_spectrogram.size(-1) * self.config.hop_len
@@ -301,4 +310,4 @@ def __call__(self, mel_spectrogram, generator):
301310
# 4. set current audio to prev_audio: x_t -> x_t-1
302311
audio = pred_prev_audio + variance
303312

304-
return audio
313+
return audio

src/diffusers/pipelines/pipeline_glide.py

Lines changed: 7 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -28,12 +28,7 @@
2828
from transformers.activations import ACT2FN
2929
from transformers.modeling_outputs import BaseModelOutput, BaseModelOutputWithPooling
3030
from transformers.modeling_utils import PreTrainedModel
31-
from transformers.utils import (
32-
ModelOutput,
33-
add_start_docstrings_to_model_forward,
34-
logging,
35-
replace_return_docstrings,
36-
)
31+
from transformers.utils import ModelOutput, add_start_docstrings_to_model_forward, logging, replace_return_docstrings
3732

3833
from ..models import GLIDESuperResUNetModel, GLIDETextToImageUNetModel
3934
from ..pipeline_utils import DiffusionPipeline
@@ -871,7 +866,12 @@ def text_model_fn(x_t, ts, transformer_out, **kwargs):
871866

872867
# Sample gaussian noise to begin loop
873868
image = torch.randn(
874-
(batch_size, self.upscale_unet.in_channels // 2, self.upscale_unet.resolution, self.upscale_unet.resolution),
869+
(
870+
batch_size,
871+
self.upscale_unet.in_channels // 2,
872+
self.upscale_unet.resolution,
873+
self.upscale_unet.resolution,
874+
),
875875
generator=generator,
876876
)
877877
image = image.to(torch_device) * upsample_temp

src/diffusers/schedulers/scheduling_ddim.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -39,7 +39,7 @@ def __init__(
3939
beta_schedule=beta_schedule,
4040
)
4141
self.timesteps = int(timesteps)
42-
self.timestep_values = timestep_values # save the fixed timestep values for BDDM
42+
self.timestep_values = timestep_values # save the fixed timestep values for BDDM
4343
self.clip_image = clip_predicted_image
4444

4545
if trained_betas is not None:

src/diffusers/schedulers/scheduling_ddpm.py

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -56,6 +56,8 @@ def __init__(
5656

5757
self.alphas = 1.0 - self.betas
5858
self.alphas_cumprod = np.cumprod(self.alphas, axis=0)
59+
self.sqrt_alphas_cumprod = np.sqrt(self.alphas_cumprod)
60+
self.sqrt_one_minus_alphas_cumprod = np.sqrt(1 - self.alphas_cumprod)
5961
self.one = np.array(1.0)
6062

6163
self.set_format(tensor_format=tensor_format)
@@ -131,5 +133,9 @@ def step(self, residual, image, t):
131133

132134
return pred_prev_image
133135

136+
def forward_step(self, original_image, noise, t):
137+
noisy_image = self.sqrt_alphas_cumprod[t] * original_image + self.sqrt_one_minus_alphas_cumprod[t] * noise
138+
return noisy_image
139+
134140
def __len__(self):
135141
return self.timesteps
Lines changed: 116 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,116 @@
1+
import random
2+
3+
import numpy as np
4+
import torch
5+
import torch.nn.functional as F
6+
7+
import PIL.Image
8+
from accelerate import Accelerator
9+
from datasets import load_dataset
10+
from diffusers import DDPM, DDPMScheduler, UNetModel
11+
from torchvision.transforms import CenterCrop, Compose, Lambda, RandomHorizontalFlip, Resize, ToTensor
12+
from tqdm.auto import tqdm
13+
from transformers import get_linear_schedule_with_warmup
14+
15+
16+
def set_seed(seed):
17+
torch.backends.cudnn.deterministic = True
18+
torch.backends.cudnn.benchmark = False
19+
torch.manual_seed(seed)
20+
torch.cuda.manual_seed_all(seed)
21+
np.random.seed(seed)
22+
random.seed(seed)
23+
24+
25+
set_seed(0)
26+
27+
accelerator = Accelerator(mixed_precision="fp16")
28+
29+
model = UNetModel(ch=128, ch_mult=(1, 2, 4, 8), resolution=64)
30+
noise_scheduler = DDPMScheduler(timesteps=1000)
31+
optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4)
32+
33+
num_epochs = 100
34+
batch_size = 8
35+
gradient_accumulation_steps = 8
36+
37+
augmentations = Compose(
38+
[
39+
Resize(64),
40+
CenterCrop(64),
41+
RandomHorizontalFlip(),
42+
ToTensor(),
43+
Lambda(lambda x: x * 2 - 1),
44+
]
45+
)
46+
dataset = load_dataset("huggan/pokemon", split="train")
47+
48+
49+
def transforms(examples):
50+
images = [augmentations(image.convert("RGB")) for image in examples["image"]]
51+
return {"input": images}
52+
53+
54+
dataset = dataset.shuffle(seed=0)
55+
dataset.set_transform(transforms)
56+
train_dataloader = torch.utils.data.DataLoader(dataset, batch_size=batch_size, shuffle=False)
57+
58+
lr_scheduler = get_linear_schedule_with_warmup(
59+
optimizer=optimizer,
60+
num_warmup_steps=1000,
61+
num_training_steps=(len(train_dataloader) * num_epochs) // gradient_accumulation_steps,
62+
)
63+
64+
model, optimizer, train_dataloader, lr_scheduler = accelerator.prepare(
65+
model, optimizer, train_dataloader, lr_scheduler
66+
)
67+
68+
for epoch in range(num_epochs):
69+
model.train()
70+
pbar = tqdm(total=len(train_dataloader), unit="ba")
71+
pbar.set_description(f"Epoch {epoch}")
72+
for step, batch in enumerate(train_dataloader):
73+
clean_images = batch["input"]
74+
noisy_images = torch.empty_like(clean_images)
75+
bsz = clean_images.shape[0]
76+
77+
timesteps = torch.randint(0, noise_scheduler.timesteps, (bsz,), device=clean_images.device).long()
78+
for idx in range(bsz):
79+
noise = torch.randn_like(clean_images[0]).to(clean_images.device)
80+
noisy_images[idx] = noise_scheduler.forward_step(clean_images[idx], noise, timesteps[idx])
81+
82+
if step % gradient_accumulation_steps == 0:
83+
with accelerator.no_sync(model):
84+
output = model(noisy_images, timesteps)
85+
loss = F.l1_loss(output, clean_images)
86+
accelerator.backward(loss)
87+
else:
88+
output = model(noisy_images, timesteps)
89+
loss = F.l1_loss(output, clean_images)
90+
accelerator.backward(loss)
91+
optimizer.step()
92+
lr_scheduler.step()
93+
optimizer.zero_grad()
94+
pbar.update(1)
95+
pbar.set_postfix(loss=loss.detach().item(), lr=optimizer.param_groups[0]["lr"])
96+
97+
optimizer.step()
98+
99+
# eval
100+
model.eval()
101+
with torch.no_grad():
102+
pipeline = DDPM(unet=model, noise_scheduler=noise_scheduler)
103+
generator = torch.Generator()
104+
generator = generator.manual_seed(0)
105+
# run pipeline in inference (sample random noise and denoise)
106+
image = pipeline(generator=generator)
107+
108+
# process image to PIL
109+
image_processed = image.cpu().permute(0, 2, 3, 1)
110+
image_processed = (image_processed + 1.0) * 127.5
111+
image_processed = image_processed.type(torch.uint8).numpy()
112+
image_pil = PIL.Image.fromarray(image_processed[0])
113+
114+
# save image
115+
pipeline.save_pretrained("./poke-ddpm")
116+
image_pil.save(f"./poke-ddpm/test_{epoch}.png")

0 commit comments

Comments
 (0)