From 9a230d742a67772c40882fdf424d31bc28a8adee Mon Sep 17 00:00:00 2001 From: Alexandru Papiu Date: Mon, 12 Feb 2024 18:50:37 +0000 Subject: [PATCH 1/7] experiment with upscaling pos encs --- tld/denoiser.py | 24 ++++++++++++++++++++++-- 1 file changed, 22 insertions(+), 2 deletions(-) diff --git a/tld/denoiser.py b/tld/denoiser.py index ac29dae..4c946de 100644 --- a/tld/denoiser.py +++ b/tld/denoiser.py @@ -20,8 +20,19 @@ def __init__(self, patch_size, img_size, embed_dim, dropout, n_layers, mlp_multi self.mlp_multiplier = mlp_multiplier seq_len = int((self.img_size/self.patch_size)*((self.img_size/self.patch_size))) + self.seq_len = seq_len patch_dim = self.n_channels*self.patch_size*self.patch_size + + self.pos_enc_down_sampling = nn.Sequential(Rearrange('bs (h w) d -> bs d h w', h=seq_len, w=seq_len), + nn.AvgPool2d(kernel_size=2), + Rearrange('bs d h w -> bs (h w) d')) + + self.pos_enc_upsampling = nn.Sequential(Rearrange('bs (h w) d -> bs d h w', h=seq_len, w=seq_len), + nn.Upsample(size=2, mode='bicubic'), + Rearrange('bs d h w -> bs (h w) d')) + + self.patchify_and_embed = nn.Sequential( nn.Conv2d(self.n_channels, patch_dim, kernel_size=self.patch_size, stride=self.patch_size), Rearrange('bs d h w -> bs (h w) d'), @@ -53,8 +64,17 @@ def __init__(self, patch_size, img_size, embed_dim, dropout, n_layers, mlp_multi def forward(self, x, cond): x = self.patchify_and_embed(x) - pos_enc = self.precomputed_pos_enc[:x.size(1)].expand(x.size(0), -1) - x = x+self.pos_embed(pos_enc) + #pos_enc = self.precomputed_pos_enc[:x.size(1)].expand(x.size(0), -1) + pos_enc = self.precomputed_pos_enc.expand(x.size(0), -1) ##bs, 256, embed_dim + + if self.seq_len > x.size(1): + ##-> embed_dim, 16, 16 -> down/upsample: + #downsample_size = self.seq_len//x.size(1) + pos_enc = self.pos_enc_down_sampling(pos_enc) + elif self.seq_len < x.size(1): + pos_enc = self.pos_enc_upsampling(pos_enc) + + x = x+self.pos_embed(pos_enc) for block in self.decoder_blocks: x = block(x, cond) From d0f12d342ca91a7f781996e244f83f0f89c1ab61 Mon Sep 17 00:00:00 2001 From: Alexandru Papiu Date: Tue, 13 Feb 2024 01:57:21 +0000 Subject: [PATCH 2/7] fix upscaling data --- tld/denoiser.py | 22 +++++++++++++--------- 1 file changed, 13 insertions(+), 9 deletions(-) diff --git a/tld/denoiser.py b/tld/denoiser.py index 4c946de..9c1ddf9 100644 --- a/tld/denoiser.py +++ b/tld/denoiser.py @@ -8,28 +8,30 @@ class DenoiserTransBlock(nn.Module): - def __init__(self, patch_size, img_size, embed_dim, dropout, n_layers, mlp_multiplier=4, n_channels=4): + def __init__(self, patch_size, img_size, embed_dim, dropout, n_layers, mlp_multiplier=4, n_channels=4, scale_factor=2): super().__init__() self.patch_size = patch_size - self.img_size = img_size + self.img_size = img_size ##size the model was trained on -> output is this * scale factor self.n_channels = n_channels self.embed_dim = embed_dim self.dropout = dropout self.n_layers = n_layers self.mlp_multiplier = mlp_multiplier + self.scale_factor = scale_factor seq_len = int((self.img_size/self.patch_size)*((self.img_size/self.patch_size))) + lat_h = lat_w = int(self.img_size/self.patch_size) self.seq_len = seq_len patch_dim = self.n_channels*self.patch_size*self.patch_size - self.pos_enc_down_sampling = nn.Sequential(Rearrange('bs (h w) d -> bs d h w', h=seq_len, w=seq_len), - nn.AvgPool2d(kernel_size=2), + self.pos_enc_down_sampling = nn.Sequential(Rearrange('bs (h w) d -> bs d h w', h=lat_h, w=lat_w), + nn.AvgPool2d(kernel_size=self.upscale_factor), Rearrange('bs d h w -> bs (h w) d')) - self.pos_enc_upsampling = nn.Sequential(Rearrange('bs (h w) d -> bs d h w', h=seq_len, w=seq_len), - nn.Upsample(size=2, mode='bicubic'), + self.pos_enc_upsampling = nn.Sequential(Rearrange('bs (h w) d -> bs d h w', h=lat_h, w=lat_w), + nn.Upsample(scale_factor=self.upscale_factor, mode='bicubic'), Rearrange('bs d h w -> bs (h w) d')) @@ -42,7 +44,8 @@ def __init__(self, patch_size, img_size, embed_dim, dropout, n_layers, mlp_multi ) self.rearrange2 = Rearrange('b (h w) (c p1 p2) -> b c (h p1) (w p2)', - h=int(self.img_size/self.patch_size), + h=int(self.img_size/self.patch_size)*self.scale_factor, + w=int(self.img_size/self.patch_size)*self.scale_factor, p1=self.patch_size, p2=self.patch_size) @@ -65,7 +68,8 @@ def __init__(self, patch_size, img_size, embed_dim, dropout, n_layers, mlp_multi def forward(self, x, cond): x = self.patchify_and_embed(x) #pos_enc = self.precomputed_pos_enc[:x.size(1)].expand(x.size(0), -1) - pos_enc = self.precomputed_pos_enc.expand(x.size(0), -1) ##bs, 256, embed_dim + pos_enc = self.precomputed_pos_enc.expand(x.size(0), -1) + pos_enc = self.pos_embed(pos_enc) ##bs, 256, embed_dim if self.seq_len > x.size(1): ##-> embed_dim, 16, 16 -> down/upsample: @@ -74,7 +78,7 @@ def forward(self, x, cond): elif self.seq_len < x.size(1): pos_enc = self.pos_enc_upsampling(pos_enc) - x = x+self.pos_embed(pos_enc) + x = x+pos_enc for block in self.decoder_blocks: x = block(x, cond) From 1f29eefb1b5924faa3765e2b1a9d60f36833196d Mon Sep 17 00:00:00 2001 From: Alexandru Papiu Date: Tue, 13 Feb 2024 02:47:58 +0000 Subject: [PATCH 3/7] fix upscaling name --- tld/denoiser.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/tld/denoiser.py b/tld/denoiser.py index 9c1ddf9..c880202 100644 --- a/tld/denoiser.py +++ b/tld/denoiser.py @@ -27,11 +27,11 @@ def __init__(self, patch_size, img_size, embed_dim, dropout, n_layers, mlp_multi self.pos_enc_down_sampling = nn.Sequential(Rearrange('bs (h w) d -> bs d h w', h=lat_h, w=lat_w), - nn.AvgPool2d(kernel_size=self.upscale_factor), + nn.AvgPool2d(kernel_size=self.scale_factor), Rearrange('bs d h w -> bs (h w) d')) self.pos_enc_upsampling = nn.Sequential(Rearrange('bs (h w) d -> bs d h w', h=lat_h, w=lat_w), - nn.Upsample(scale_factor=self.upscale_factor, mode='bicubic'), + nn.Upsample(scale_factor=self.scale_factor, mode='bicubic'), Rearrange('bs d h w -> bs (h w) d')) From c593ddd473f39d006c91faf3b9f15d736903b72b Mon Sep 17 00:00:00 2001 From: Alexandru Papiu Date: Tue, 13 Feb 2024 16:15:43 +0000 Subject: [PATCH 4/7] compile model before training --- tld/train.py | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/tld/train.py b/tld/train.py index 5d665ab..c0e0226 100644 --- a/tld/train.py +++ b/tld/train.py @@ -28,6 +28,7 @@ def eval_gen(diffuser, labels): num_imgs=64, class_guidance=class_guidance, seed=seed, + img_size=64, ###hardcode bad n_iter=40, exponent=1, sharp_f=0.1, @@ -110,6 +111,9 @@ def main(config: ModelConfig, dataconfig: DataConfig): loss_fn = nn.MSELoss() optimizer = torch.optim.Adam(model.parameters(), lr=config.lr) + accelerator.print("Compiling model:") + model = torch.compile(model) + if not config.from_scratch: accelerator.print("Loading Model:") wandb.restore(config.model_name, run_path=f"apapiu/cifar_diffusion/runs/{config.run_id}", From 0c3d4500728637fa71ac19e923fbb287125ee5ea Mon Sep 17 00:00:00 2001 From: Alexandru Papiu Date: Tue, 13 Feb 2024 18:45:48 +0000 Subject: [PATCH 5/7] change to bilinear, bicubic too expensive --- tld/denoiser.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tld/denoiser.py b/tld/denoiser.py index c880202..a974a34 100644 --- a/tld/denoiser.py +++ b/tld/denoiser.py @@ -31,7 +31,7 @@ def __init__(self, patch_size, img_size, embed_dim, dropout, n_layers, mlp_multi Rearrange('bs d h w -> bs (h w) d')) self.pos_enc_upsampling = nn.Sequential(Rearrange('bs (h w) d -> bs d h w', h=lat_h, w=lat_w), - nn.Upsample(scale_factor=self.scale_factor, mode='bicubic'), + nn.Upsample(scale_factor=self.scale_factor, mode='bilinear'), Rearrange('bs d h w -> bs (h w) d')) From 024e02fc4494f201218609871845ef7458bc1ee4 Mon Sep 17 00:00:00 2001 From: Alexandru Papiu Date: Tue, 13 Feb 2024 18:48:29 +0000 Subject: [PATCH 6/7] change eval defaults --- tld/train.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/tld/train.py b/tld/train.py index c0e0226..bec87b2 100644 --- a/tld/train.py +++ b/tld/train.py @@ -24,17 +24,17 @@ def eval_gen(diffuser, labels): class_guidance=4.5 seed=10 - out, _ = diffuser.generate(labels=torch.repeat_interleave(labels, 8, dim=0), - num_imgs=64, + out, _ = diffuser.generate(labels=torch.repeat_interleave(labels, 4, dim=0), + num_imgs=32, class_guidance=class_guidance, seed=seed, img_size=64, ###hardcode bad - n_iter=40, + n_iter=20, exponent=1, sharp_f=0.1, ) - out = to_pil((vutils.make_grid((out+1)/2, nrow=8, padding=4)).float().clip(0, 1)) + out = to_pil((vutils.make_grid((out+1)/2, nrow=4, padding=4)).float().clip(0, 1)) out.save(f'emb_val_cfg:{class_guidance}_seed:{seed}.png') return out From c175eebc8cf391a4937669a42a5fc73dd400bbb8 Mon Sep 17 00:00:00 2001 From: Alexandru Papiu Date: Sat, 17 Feb 2024 05:45:59 +0000 Subject: [PATCH 7/7] improve test and parametrize scale factor --- tld/denoiser.py | 36 ++++++++++++++++++++++++++++-------- 1 file changed, 28 insertions(+), 8 deletions(-) diff --git a/tld/denoiser.py b/tld/denoiser.py index a974a34..61e4845 100644 --- a/tld/denoiser.py +++ b/tld/denoiser.py @@ -8,7 +8,7 @@ class DenoiserTransBlock(nn.Module): - def __init__(self, patch_size, img_size, embed_dim, dropout, n_layers, mlp_multiplier=4, n_channels=4, scale_factor=2): + def __init__(self, patch_size, img_size, embed_dim, dropout, n_layers, mlp_multiplier=4, n_channels=4, scale_factor=1): super().__init__() self.patch_size = patch_size @@ -88,7 +88,7 @@ def forward(self, x, cond): class Denoiser(nn.Module): def __init__(self, image_size, noise_embed_dims, patch_size, embed_dim, dropout, n_layers, - text_emb_size=768): + text_emb_size=768, scale_factor=1): super().__init__() self.image_size = image_size @@ -101,7 +101,7 @@ def __init__(self, nn.Linear(self.embed_dim, self.embed_dim) ) - self.denoiser_trans_block = DenoiserTransBlock(patch_size, image_size, embed_dim, dropout, n_layers) + self.denoiser_trans_block = DenoiserTransBlock(patch_size, image_size, embed_dim, dropout, n_layers, scale_factor=scale_factor) self.norm = nn.LayerNorm(self.embed_dim) self.label_proj = nn.Linear(text_emb_size, self.embed_dim) @@ -118,14 +118,34 @@ def forward(self, x, noise_level, label): return x -def test_outputs(): - model = Denoiser(image_size=16, noise_embed_dims=128, patch_size=2, embed_dim=256, dropout=0.1, n_layers=6) - x = torch.rand(8, 4, 16, 16) +def test_outputs(num_imgs = 1): + import time + + model = Denoiser(image_size=32, noise_embed_dims=128, patch_size=2, embed_dim=768, dropout=0.1, n_layers=12) + x = torch.rand(num_imgs, 4, 32, 32) + noise_level = torch.rand(num_imgs, 1) + label = torch.rand(num_imgs, 768) + + print(f"Model has {sum(p.numel() for p in model.parameters())} parameters") + + with torch.no_grad(): + start_time = time.time() + output = model(x, noise_level, label) + end_time = time.time() + + execution_time = end_time - start_time + print(f"Model execution took {execution_time:.4f} seconds.") + + assert output.shape == torch.Size([num_imgs, 4, 32, 32]) + print("Basic tests passed.") + + model = Denoiser(image_size=16, noise_embed_dims=128, patch_size=2, embed_dim=256, dropout=0.1, n_layers=6, scale_factor=2) + x = torch.rand(8, 4, 32, 32) noise_level = torch.rand(8, 1) label = torch.rand(8, 768) with torch.no_grad(): output = model(x, noise_level, label) - assert output.shape == torch.Size([8, 4, 16, 16]) - print("Basic tests passed.") \ No newline at end of file + assert output.shape == torch.Size([8, 4, 32, 32]) + print("Uspscale tests passed.") \ No newline at end of file