From 908c57cbd0558657530083f333de15933fb9d4dc Mon Sep 17 00:00:00 2001 From: Alexandru Papiu Date: Sun, 4 Feb 2024 23:00:04 +0000 Subject: [PATCH 1/6] use pytorch instead of rearrange and add compile step --- tld/train.py | 4 ++++ tld/transformer_blocks.py | 16 ++++++++++++---- 2 files changed, 16 insertions(+), 4 deletions(-) diff --git a/tld/train.py b/tld/train.py index 5d665ab..b5cee21 100644 --- a/tld/train.py +++ b/tld/train.py @@ -107,9 +107,13 @@ def main(config: ModelConfig, dataconfig: DataConfig): patch_size=config.patch_size, embed_dim=config.embed_dim, dropout=config.dropout, n_layers=config.n_layers) + 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}", diff --git a/tld/transformer_blocks.py b/tld/transformer_blocks.py index 076ce74..fb3f1b3 100644 --- a/tld/transformer_blocks.py +++ b/tld/transformer_blocks.py @@ -30,15 +30,23 @@ def forward(self, q, k, v, attn_mask=None): assert q.size(-1) == k.size(-1) assert k.size(-2) == v.size(-2) - q, k, v = [rearrange(x, 'bs n (d h) -> bs h n d', h=self.n_heads) for x in [q,k,v]] - q, k, v = q.contiguous(), k.contiguous(), v.contiguous() + bs = q.size(0) + d_k = q.size(-1) // self.n_heads + + q = q.view(bs, q.size(1), self.n_heads, d_k).permute(0, 2, 1, 3) + k = k.view(bs, k.size(1), self.n_heads, d_k).permute(0, 2, 1, 3) + v = v.view(bs, v.size(1), self.n_heads, d_k).permute(0, 2, 1, 3) + + # q, k, v = [rearrange(x, 'bs n (d h) -> bs h n d', h=self.n_heads) for x in [q,k,v]] + # q, k, v = q.contiguous(), k.contiguous(), v.contiguous() out = nn.functional.scaled_dot_product_attention(q, k, v, attn_mask=attn_mask, is_causal=self.is_causal, dropout_p=self.dropout_level if self.training else 0) - - out = rearrange(out, 'bs h n d -> bs n (d h)', h=self.n_heads) + + out = out.permute(0, 2, 1, 3).contiguous().view(bs, out.size(1), -1) + #out = rearrange(out, 'bs h n d -> bs n (d h)', h=self.n_heads) return out From 4a698f672ae061a9d2b052f02659155adab08ce6 Mon Sep 17 00:00:00 2001 From: Alexandru Papiu Date: Sun, 4 Feb 2024 23:28:55 +0000 Subject: [PATCH 2/6] reshape fix --- tld/transformer_blocks.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/tld/transformer_blocks.py b/tld/transformer_blocks.py index fb3f1b3..7630bfc 100644 --- a/tld/transformer_blocks.py +++ b/tld/transformer_blocks.py @@ -30,7 +30,7 @@ def forward(self, q, k, v, attn_mask=None): assert q.size(-1) == k.size(-1) assert k.size(-2) == v.size(-2) - bs = q.size(0) + bs, seq_len, _ = q.size(0) d_k = q.size(-1) // self.n_heads q = q.view(bs, q.size(1), self.n_heads, d_k).permute(0, 2, 1, 3) @@ -45,7 +45,7 @@ def forward(self, q, k, v, attn_mask=None): is_causal=self.is_causal, dropout_p=self.dropout_level if self.training else 0) - out = out.permute(0, 2, 1, 3).contiguous().view(bs, out.size(1), -1) + out = out.permute(0, 2, 1, 3).contiguous().view(bs, seq_len, -1) #out = rearrange(out, 'bs h n d -> bs n (d h)', h=self.n_heads) return out From f8f636e02e5c5686eabdbeaa4a95cfe5d565afc1 Mon Sep 17 00:00:00 2001 From: Alexandru Papiu Date: Mon, 5 Feb 2024 02:44:44 +0000 Subject: [PATCH 3/6] change the reshape to correctly put head first --- tld/transformer_blocks.py | 13 ++----------- 1 file changed, 2 insertions(+), 11 deletions(-) diff --git a/tld/transformer_blocks.py b/tld/transformer_blocks.py index 7630bfc..5de41fa 100644 --- a/tld/transformer_blocks.py +++ b/tld/transformer_blocks.py @@ -30,23 +30,14 @@ def forward(self, q, k, v, attn_mask=None): assert q.size(-1) == k.size(-1) assert k.size(-2) == v.size(-2) - bs, seq_len, _ = q.size(0) - d_k = q.size(-1) // self.n_heads - - q = q.view(bs, q.size(1), self.n_heads, d_k).permute(0, 2, 1, 3) - k = k.view(bs, k.size(1), self.n_heads, d_k).permute(0, 2, 1, 3) - v = v.view(bs, v.size(1), self.n_heads, d_k).permute(0, 2, 1, 3) - - # q, k, v = [rearrange(x, 'bs n (d h) -> bs h n d', h=self.n_heads) for x in [q,k,v]] - # q, k, v = q.contiguous(), k.contiguous(), v.contiguous() + q, k, v = [rearrange(x, 'bs n (h d) -> bs h n d', h=self.n_heads) for x in [q,k,v]] out = nn.functional.scaled_dot_product_attention(q, k, v, attn_mask=attn_mask, is_causal=self.is_causal, dropout_p=self.dropout_level if self.training else 0) - out = out.permute(0, 2, 1, 3).contiguous().view(bs, seq_len, -1) - #out = rearrange(out, 'bs h n d -> bs n (d h)', h=self.n_heads) + out = rearrange(out, 'bs h n d -> bs n (h d)', h=self.n_heads) return out From b249d8a9faaa42624446312f368c9b37e03b629b Mon Sep 17 00:00:00 2001 From: Alexandru Papiu Date: Mon, 5 Feb 2024 04:06:42 +0000 Subject: [PATCH 4/6] add pre norm --- tld/transformer_blocks.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/tld/transformer_blocks.py b/tld/transformer_blocks.py index 5de41fa..2f87ddb 100644 --- a/tld/transformer_blocks.py +++ b/tld/transformer_blocks.py @@ -101,14 +101,14 @@ class DecoderBlock(nn.Module): def __init__(self, embed_dim, is_causal, mlp_multiplier, dropout_level, mlp_class=MLP): super().__init__() self.self_attention = SelfAttention(embed_dim, is_causal, dropout_level, n_heads=embed_dim//64) - self.cross_attention = CrossAttention(embed_dim, is_causal=False, dropout_level=0, n_heads=4) + self.cross_attention = CrossAttention(embed_dim, is_causal=False, dropout_level=0, n_heads=embed_dim//64) self.mlp = mlp_class(embed_dim, mlp_multiplier, dropout_level) self.norm1 = nn.LayerNorm(embed_dim) self.norm2 = nn.LayerNorm(embed_dim) self.norm3 = nn.LayerNorm(embed_dim) def forward(self, x, y): - x = self.norm1(self.self_attention(x) + x) - x = self.norm2(self.cross_attention(x, y) + x) - x = self.norm3(self.mlp(x) + x) + x = self.self_attention(self.norm1(x)) + x + x = self.cross_attention(self.norm2(x), y) + x + x = self.mlp(self.norm3(x)) + x return x \ No newline at end of file From d3721b3d06c408d8928d65c528bfe97d08027078 Mon Sep 17 00:00:00 2001 From: Alexandru Papiu Date: Tue, 6 Feb 2024 22:13:00 +0000 Subject: [PATCH 5/6] new architecture for attention at highest rez --- tld/denoiser.py | 32 ++++++++++++++++++++++++++------ 1 file changed, 26 insertions(+), 6 deletions(-) diff --git a/tld/denoiser.py b/tld/denoiser.py index ac29dae..48ec0ac 100644 --- a/tld/denoiser.py +++ b/tld/denoiser.py @@ -19,14 +19,22 @@ def __init__(self, patch_size, img_size, embed_dim, dropout, n_layers, mlp_multi self.n_layers = n_layers self.mlp_multiplier = mlp_multiplier - seq_len = int((self.img_size/self.patch_size)*((self.img_size/self.patch_size))) + seq_len = img_size**2#int((self.img_size/self.patch_size)*((self.img_size/self.patch_size))) patch_dim = self.n_channels*self.patch_size*self.patch_size + self.first_proj = nn.Sequential(Rearrange('bs d h w -> bs (h w) d'), + nn.Linear(n_channels, self.embed_dim//4), + nn.LayerNorm(self.embed_dim//4), + ) + + self.cond_proj = nn.Linear(self.embed_dim, self.embed_dim//4) + self.patchify_and_embed = nn.Sequential( - nn.Conv2d(self.n_channels, patch_dim, kernel_size=self.patch_size, stride=self.patch_size), + Rearrange('bs (h w) d ->bs d h w', h = self.img_size), + nn.Conv2d(self.embed_dim//4, self.embed_dim, kernel_size=self.patch_size, stride=self.patch_size), Rearrange('bs d h w -> bs (h w) d'), - nn.LayerNorm(patch_dim), - nn.Linear(patch_dim, self.embed_dim), + nn.LayerNorm(self.embed_dim), + nn.Linear(self.embed_dim, self.embed_dim), nn.LayerNorm(self.embed_dim) ) @@ -35,9 +43,15 @@ def __init__(self, patch_size, img_size, embed_dim, dropout, n_layers, mlp_multi p1=self.patch_size, p2=self.patch_size) - self.pos_embed = nn.Embedding(seq_len, self.embed_dim) + self.pos_embed = nn.Embedding(seq_len, self.embed_dim//4) self.register_buffer('precomputed_pos_enc', torch.arange(0, seq_len).long()) + self.first_decoder_block = DecoderBlock(embed_dim=self.embed_dim//4, ##aplied on priginal image + mlp_multiplier=self.mlp_multiplier, + is_causal=False, + dropout_level=self.dropout, + mlp_class=MLPSepConv) + self.decoder_blocks = nn.ModuleList([DecoderBlock(embed_dim=self.embed_dim, mlp_multiplier=self.mlp_multiplier, #note that this is a non-causal block since we are @@ -52,10 +66,16 @@ 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) + x = self.first_proj(x) ### 4,h,w -> h*w, 4 -> h*w, d/4 + pos_enc = self.precomputed_pos_enc[:x.size(1)].expand(x.size(0), -1) x = x+self.pos_embed(pos_enc) + ##have to project cond: + cond_projected = self.cond_proj(cond) + x = self.first_decoder_block(x, cond_projected) ###same dim as original data in and out. + x = self.patchify_and_embed(x) ## h*w, d/4 -> d/4, h, w -> h/2*w/2, d + for block in self.decoder_blocks: x = block(x, cond) From 5e5e389499250a6542aa5fca40d61145615dc64b Mon Sep 17 00:00:00 2001 From: Alexandru Papiu Date: Fri, 9 Feb 2024 23:57:52 +0000 Subject: [PATCH 6/6] comments for later --- tld/denoiser.py | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/tld/denoiser.py b/tld/denoiser.py index 48ec0ac..11310e3 100644 --- a/tld/denoiser.py +++ b/tld/denoiser.py @@ -79,6 +79,11 @@ def forward(self, x, cond): for block in self.decoder_blocks: x = block(x, cond) + ## should I add another attention layer here at 1024 tokens? + + #h/2*w/2, d -> d, h/2, w/2 -> upsample -> attenion block -> h*w, d -> h*w, 4 -> + + return self.out_proj(x) class Denoiser(nn.Module):