Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
68 changes: 56 additions & 12 deletions tld/denoiser.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,20 +8,33 @@


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=1):
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=lat_h, w=lat_w),
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.scale_factor, mode='bilinear'),
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'),
Expand All @@ -31,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)


Expand All @@ -53,8 +67,18 @@ 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)
pos_enc = self.pos_embed(pos_enc) ##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+pos_enc

for block in self.decoder_blocks:
x = block(x, cond)
Expand All @@ -64,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
Expand All @@ -77,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)

Expand All @@ -94,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.")
assert output.shape == torch.Size([8, 4, 32, 32])
print("Uspscale tests passed.")
12 changes: 8 additions & 4 deletions tld/train.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,16 +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,
n_iter=40,
img_size=64, ###hardcode bad
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
Expand Down Expand Up @@ -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}",
Expand Down