Rohan3 commited on
Commit
a625e96
·
1 Parent(s): 87b5061

Updated: VAE, UNet, config, text embeddings, model and main

Browse files
app/main.py CHANGED
@@ -36,9 +36,10 @@ app.add_middleware(
36
  class GenerateRequest(BaseModel):
37
  caption: str = Field(..., example="a white dog running in snow")
38
  num_images: int = Field(4, ge=1, le=8)
39
- num_steps: int = Field(50, ge=10, le=200)
40
- guidance_scale: float = Field(7.5, ge=1.0, le=20.0)
41
  seed: int = Field(42)
 
42
 
43
  def tensor_to_pil(img_tensor: torch.Tensor) -> Image.Image:
44
  img = img_tensor.clamp(0, 1)
@@ -78,83 +79,9 @@ async def generate(req: GenerateRequest, _=Depends(verify_key)):
78
  num_steps=req.num_steps,
79
  guidance_scale=req.guidance_scale,
80
  seed=req.seed,
 
81
  )
82
  b64_images = [tensor_to_base64(img) for img in images]
83
  return GenerateResponse(images=b64_images, num_generated=len(b64_images))
84
- except Exception as e:
85
- raise HTTPException(status_code=500, detail=str(e))
86
-
87
- # ── Single image (num_images=1) ──────────────────────────────────────
88
- @app.post("/generate/image", response_class=StreamingResponse)
89
- async def generate_image(req: GenerateRequest):
90
- try:
91
- req.num_images = 1
92
- images = pipeline.generate(
93
- caption=req.caption,
94
- num_images=1,
95
- num_steps=req.num_steps,
96
- guidance_scale=req.guidance_scale,
97
- seed=req.seed,
98
- )
99
- buf = io.BytesIO()
100
- tensor_to_pil(images[0]).save(buf, format="PNG")
101
- buf.seek(0)
102
- return StreamingResponse(buf, media_type="image/png")
103
- except Exception as e:
104
- raise HTTPException(status_code=500, detail=str(e))
105
-
106
- # ── Multiple images as ZIP ───────────────────────────────────────────
107
- @app.post("/generate/zip", response_class=StreamingResponse)
108
- async def generate_zip(req: GenerateRequest):
109
- try:
110
- images = pipeline.generate(
111
- caption=req.caption,
112
- num_images=req.num_images,
113
- num_steps=req.num_steps,
114
- guidance_scale=req.guidance_scale,
115
- seed=req.seed,
116
- )
117
- buf = io.BytesIO()
118
- with zipfile.ZipFile(buf, "w") as zf:
119
- for i, img_tensor in enumerate(images):
120
- img_buf = io.BytesIO()
121
- tensor_to_pil(img_tensor).save(img_buf, format="PNG")
122
- zf.writestr(f"image_{i}.png", img_buf.getvalue())
123
- buf.seek(0)
124
- return StreamingResponse(buf, media_type="application/zip",
125
- headers={"Content-Disposition": "attachment; filename=images.zip"})
126
- except Exception as e:
127
- raise HTTPException(status_code=500, detail=str(e))
128
-
129
- # ── Preview page: see all images directly in browser ─────────────────
130
- @app.post("/generate/preview", response_class=HTMLResponse)
131
- async def generate_preview(req: GenerateRequest):
132
- try:
133
- images = pipeline.generate(
134
- caption=req.caption,
135
- num_images=req.num_images,
136
- num_steps=req.num_steps,
137
- guidance_scale=req.guidance_scale,
138
- seed=req.seed,
139
- )
140
- img_tags = ""
141
- for img_tensor in images:
142
- transform = T.ToPILImage()
143
- img = transform(img_tensor)
144
- img.show()
145
- buf = io.BytesIO()
146
- tensor_to_pil(img_tensor).save(buf, format="PNG")
147
- b64 = base64.b64encode(buf.getvalue()).decode("utf-8")
148
- img_tags += f'<img src="data:image/png;base64,{b64}" style="width:256px; margin:5px;"/>'
149
-
150
- html = f"""
151
- <html>
152
- <body style="background:#111; display:flex; flex-wrap:wrap;">
153
- <h2 style="color:white; width:100%;">Caption: {req.caption}</h2>
154
- {img_tags}
155
- </body>
156
- </html>
157
- """
158
- return HTMLResponse(content=html)
159
  except Exception as e:
160
  raise HTTPException(status_code=500, detail=str(e))
 
36
  class GenerateRequest(BaseModel):
37
  caption: str = Field(..., example="a white dog running in snow")
38
  num_images: int = Field(4, ge=1, le=8)
39
+ num_steps: int = Field(30, ge=10, le=100)
40
+ guidance_scale: float = Field(5, ge=1.0, le=20.0)
41
  seed: int = Field(42)
42
+ eta: float = Field(0, ge=0.0, le=1.0)
43
 
44
  def tensor_to_pil(img_tensor: torch.Tensor) -> Image.Image:
45
  img = img_tensor.clamp(0, 1)
 
79
  num_steps=req.num_steps,
80
  guidance_scale=req.guidance_scale,
81
  seed=req.seed,
82
+ eta=req.eta,
83
  )
84
  b64_images = [tensor_to_base64(img) for img in images]
85
  return GenerateResponse(images=b64_images, num_generated=len(b64_images))
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
86
  except Exception as e:
87
  raise HTTPException(status_code=500, detail=str(e))
app/model.py CHANGED
@@ -84,14 +84,12 @@ class LDMPipeline:
84
  per_token_contextual = self.text_model.ln_final(x) # (B, T, D) = (1, 77, 1024)
85
  return per_token_contextual.squeeze(0) # (77, 1024)`
86
 
87
- def generate(self, caption: str, num_images: int = 4,
88
- num_steps: int = 50, guidance_scale: float = 7.5,
89
- seed: int = 42):
90
  seed_everything(seed)
91
  caption = caption.strip()
92
  if caption.endswith("."):
93
  caption = caption.rstrip(".")
94
- caption = caption.lower()
95
  embedding = self.get_text_embedding(caption).unsqueeze(0)
96
  latents = ddim_sample(
97
  unet=self.unet,
@@ -101,7 +99,7 @@ class LDMPipeline:
101
  embedding=embedding,
102
  guidance_scale=guidance_scale,
103
  num_steps=num_steps,
104
- eta=0.0,
105
  device=self.device
106
  )
107
  latents = latents * latent_std
 
84
  per_token_contextual = self.text_model.ln_final(x) # (B, T, D) = (1, 77, 1024)
85
  return per_token_contextual.squeeze(0) # (77, 1024)`
86
 
87
+ def generate(self, caption: str, num_images: int = 4, num_steps: int = 50, guidance_scale: float = 7.5, seed: int = 42, eta: float = 0):
 
 
88
  seed_everything(seed)
89
  caption = caption.strip()
90
  if caption.endswith("."):
91
  caption = caption.rstrip(".")
92
+ # caption = caption.lower()
93
  embedding = self.get_text_embedding(caption).unsqueeze(0)
94
  latents = ddim_sample(
95
  unet=self.unet,
 
99
  embedding=embedding,
100
  guidance_scale=guidance_scale,
101
  num_steps=num_steps,
102
+ eta=eta,
103
  device=self.device
104
  )
105
  latents = latents * latent_std
core/config.py CHANGED
@@ -1,26 +1,29 @@
1
  tqdm_colors = ["#0867C5", "green", "blue", "red", "yellow", "#6707B0", "#FFDD22", "#00FFEF", "#DC442C"]
2
  data_dir = "./backend/core/data"
3
- image_res = 128
4
  image_dir = f"{data_dir}/Images"
5
  resized_img_dir = f"{data_dir}/ResizedImages_{image_res}"
6
- vae_batch_size = 8
7
  vae_group_size = 4
8
  vae_num_epochs = 1000
9
  vae_stopping_patience = 30
10
  vae_latent_channels = 4 # 16, 8
11
  vae_latent_dim = 32
12
  vae_beta_kld = 1e-3
13
- vae_optim_lr = 1e-4
14
  vae_lambda_tvl = 0 # 1e-3
15
  vae_lpips_weight = 1e-1
16
  vae_dropout = 0.
17
  vae_checkpoint_dir = "./backend/core/checkpoints/vae"
18
- vae_weight = f"attn_test_best_{image_res}_3_4_32_128_{vae_latent_channels}_beta_{vae_beta_kld}_tvl_{vae_lambda_tvl}_batch_{vae_batch_size}_lr_{vae_optim_lr}.pth"
 
19
  latent_dir = f"./backend/core/latents_{image_res}"
20
  latent_scaled_dir = f"./backend/core/latents_scaled_{image_res}"
21
  latent_norm_dir = f"./backend/core/latents_norm_{image_res}"
22
  latent_recon_images = f"./backend/core/recon_img_{vae_latent_channels}_{vae_latent_dim}_{vae_latent_dim}_res_{image_res}"
23
- latent_mu = 0.021192772313952446; latent_std = 0.9767765402793884 # Latents
 
 
24
  # Mean over latent: 0.021192772313952446
25
  # STD over latent: 0.9767765402793884
26
  # Latent Scale: 1.023775577545166
@@ -39,19 +42,20 @@ text_captions_dir = f"{data_dir}/captions.txt"
39
  unet_batch_size = 32
40
  embedding_dim = 1024
41
  embedding_dir = f"./backend/core/embeddings_77_{embedding_dim}"
 
42
  embedding_model = "ViT-g-14" # or "ViT-B-16"
43
  embedding_pretrained = "laion2b_s12b_b42k" # or "openai"
44
  unet_pred_type = "v_prediction" # v_prediction or epsilon
45
  unet_checkpoint_dir = f"./backend/core/checkpoints/ldm"
46
  unet_max_steps = 256_000
47
  # new_lr = old_lr * (batch_size_new / batch_size_old)
48
- unet_optim_lr = 1e-4
49
  unet_group_size = 32
50
  unet_beta_schedule = "squaredcos_cap_v2" # "linear" or "squaredcos_cap_v2"
51
  unet_dropout = 0.
52
  attn_dropout = 0.
53
  unet_train_timesteps = 1000
54
- ddim_guidace_scale = 3
55
  ddim_num_sampling_steps = 100
56
  ddim_img_dir = "./backend/core/ddim_recon_img"
57
  unet_val_embeddings_dir = "./backend/core/embeddings_val"
 
1
  tqdm_colors = ["#0867C5", "green", "blue", "red", "yellow", "#6707B0", "#FFDD22", "#00FFEF", "#DC442C"]
2
  data_dir = "./backend/core/data"
3
+ image_res = 256
4
  image_dir = f"{data_dir}/Images"
5
  resized_img_dir = f"{data_dir}/ResizedImages_{image_res}"
6
+ vae_batch_size = 4
7
  vae_group_size = 4
8
  vae_num_epochs = 1000
9
  vae_stopping_patience = 30
10
  vae_latent_channels = 4 # 16, 8
11
  vae_latent_dim = 32
12
  vae_beta_kld = 1e-3
13
+ vae_optim_lr = 5e-5
14
  vae_lambda_tvl = 0 # 1e-3
15
  vae_lpips_weight = 1e-1
16
  vae_dropout = 0.
17
  vae_checkpoint_dir = "./backend/core/checkpoints/vae"
18
+ # vae_weight = f"attn_test_best_{image_res}_3_32_128_256_512_{vae_latent_channels}_beta_{vae_beta_kld}_tvl_{vae_lambda_tvl}_batch_{vae_batch_size}_lr_{vae_optim_lr}.pth"
19
+ vae_weight = f"ssim_lpips_attn_test_best_{image_res}_3_32_128_256_512_{vae_latent_channels}_beta_{vae_beta_kld}_tvl_{vae_lambda_tvl}_batch_{vae_batch_size}_lr_{vae_optim_lr}.pth"
20
  latent_dir = f"./backend/core/latents_{image_res}"
21
  latent_scaled_dir = f"./backend/core/latents_scaled_{image_res}"
22
  latent_norm_dir = f"./backend/core/latents_norm_{image_res}"
23
  latent_recon_images = f"./backend/core/recon_img_{vae_latent_channels}_{vae_latent_dim}_{vae_latent_dim}_res_{image_res}"
24
+ # latent_mu = 0.021192772313952446; latent_std = 0.9767765402793884 # Latents
25
+ # latent_mu = 0.0007418220047838986; latent_std = 0.9822604060173035 # Latents scale = 1.0180599689483643
26
+ latent_mu = -0.08573896437883377; latent_std = 1.2452856302261353 # Latents scale = 0.8030286431312561
27
  # Mean over latent: 0.021192772313952446
28
  # STD over latent: 0.9767765402793884
29
  # Latent Scale: 1.023775577545166
 
42
  unet_batch_size = 32
43
  embedding_dim = 1024
44
  embedding_dir = f"./backend/core/embeddings_77_{embedding_dim}"
45
+ null_embedding_dir = f"./backend/core/embeddings_77_{embedding_dim}/null_embedding.pt"
46
  embedding_model = "ViT-g-14" # or "ViT-B-16"
47
  embedding_pretrained = "laion2b_s12b_b42k" # or "openai"
48
  unet_pred_type = "v_prediction" # v_prediction or epsilon
49
  unet_checkpoint_dir = f"./backend/core/checkpoints/ldm"
50
  unet_max_steps = 256_000
51
  # new_lr = old_lr * (batch_size_new / batch_size_old)
52
+ unet_optim_lr = 1e-4 # Changed to 5e-5 after 1103080 steps
53
  unet_group_size = 32
54
  unet_beta_schedule = "squaredcos_cap_v2" # "linear" or "squaredcos_cap_v2"
55
  unet_dropout = 0.
56
  attn_dropout = 0.
57
  unet_train_timesteps = 1000
58
+ ddim_guidace_scale = 8
59
  ddim_num_sampling_steps = 100
60
  ddim_img_dir = "./backend/core/ddim_recon_img"
61
  unet_val_embeddings_dir = "./backend/core/embeddings_val"
core/dataloader.py CHANGED
@@ -21,13 +21,13 @@ class ImageDataset(Dataset):
21
  def image_dataloader():
22
  image_paths = [os.path.join(resized_img_dir, path) for path in os.listdir(resized_img_dir)[:] if path.endswith(".jpg")]
23
  dataset = ImageDataset(image_paths)
24
- g = torch.Generator()
25
- g.manual_seed(42)
26
  train_size = int(0.8 * len(dataset))
27
  test_size = len(dataset) - train_size
28
  train_set, test_set = random_split(dataset, [train_size, test_size])
29
- train_loader = DataLoader(train_set, batch_size=vae_batch_size, shuffle=True, generator=g, worker_init_fn=seed_worker, num_workers=min(8, os.cpu_count()), pin_memory=True, persistent_workers=True, prefetch_factor=16)
30
- test_loader = DataLoader(test_set, batch_size=vae_batch_size, shuffle=False, generator=g, worker_init_fn=seed_worker, num_workers=min(8, os.cpu_count()), pin_memory=True, persistent_workers=True, prefetch_factor=16)
31
  return train_loader, test_loader
32
 
33
  class LatentEmbeddingsDataset(Dataset):
@@ -39,20 +39,21 @@ class LatentEmbeddingsDataset(Dataset):
39
  def __getitem__(self, index):
40
  latent_index = index // self.num_variants
41
  variant_index = index % self.num_variants
 
42
  file_path = self.file_names[latent_index]
43
  latent_data = torch.load(f"{latent_scaled_dir}/{file_path}", map_location="cpu", weights_only=True)
44
  embedding_data = torch.load(f"{embedding_dir}/{file_path[:-3]}_{variant_index}.pt", map_location="cpu", weights_only=True)
45
  assert embedding_data.shape == torch.Size([77, 1024]), f"Unexpected embedding shape: {embedding_data.shape}"
46
  return (latent_data, embedding_data)
47
  def latent_embedding_dataloader():
48
- file_names = [path for path in os.listdir(latent_scaled_dir) if path.endswith(".pt")][:]
 
49
  dataset = LatentEmbeddingsDataset(file_names)
50
- g = torch.Generator()
51
- g.manual_seed(42)
52
- train_loader = DataLoader(dataset, batch_size=unet_batch_size, shuffle=True, generator=g, worker_init_fn=seed_worker, num_workers=min(8, os.cpu_count()), pin_memory=True, persistent_workers=True, prefetch_factor=32)
53
  return train_loader
54
 
55
  def seed_worker(worker_id):
56
- worker_seed = 42 + worker_id
 
57
  np.random.seed(worker_seed)
58
  random.seed(worker_seed)
 
21
  def image_dataloader():
22
  image_paths = [os.path.join(resized_img_dir, path) for path in os.listdir(resized_img_dir)[:] if path.endswith(".jpg")]
23
  dataset = ImageDataset(image_paths)
24
+ # g = torch.Generator()
25
+ # g.manual_seed(42)
26
  train_size = int(0.8 * len(dataset))
27
  test_size = len(dataset) - train_size
28
  train_set, test_set = random_split(dataset, [train_size, test_size])
29
+ train_loader = DataLoader(train_set, batch_size=vae_batch_size, shuffle=True, worker_init_fn=seed_worker, num_workers=min(8, os.cpu_count()), pin_memory=True, persistent_workers=True, prefetch_factor=4)
30
+ test_loader = DataLoader(test_set, batch_size=vae_batch_size, shuffle=False, worker_init_fn=seed_worker, num_workers=min(8, os.cpu_count()), pin_memory=True, persistent_workers=True, prefetch_factor=4)
31
  return train_loader, test_loader
32
 
33
  class LatentEmbeddingsDataset(Dataset):
 
39
  def __getitem__(self, index):
40
  latent_index = index // self.num_variants
41
  variant_index = index % self.num_variants
42
+ # variant_index = random.randint(0, self.num_variants - 1)
43
  file_path = self.file_names[latent_index]
44
  latent_data = torch.load(f"{latent_scaled_dir}/{file_path}", map_location="cpu", weights_only=True)
45
  embedding_data = torch.load(f"{embedding_dir}/{file_path[:-3]}_{variant_index}.pt", map_location="cpu", weights_only=True)
46
  assert embedding_data.shape == torch.Size([77, 1024]), f"Unexpected embedding shape: {embedding_data.shape}"
47
  return (latent_data, embedding_data)
48
  def latent_embedding_dataloader():
49
+ file_names = sorted(path for path in os.listdir(latent_scaled_dir) if path.endswith(".pt"))[:]
50
+ # file_names = [path for path in os.listdir(latent_scaled_dir) if path.endswith(".pt")][:16]
51
  dataset = LatentEmbeddingsDataset(file_names)
52
+ train_loader = DataLoader(dataset, batch_size=unet_batch_size, shuffle=True, worker_init_fn=seed_worker, num_workers=min(8, os.cpu_count()), pin_memory=True, persistent_workers=True, prefetch_factor=32)
 
 
53
  return train_loader
54
 
55
  def seed_worker(worker_id):
56
+ # worker_seed = 42 + worker_id
57
+ worker_seed = torch.initial_seed() % 2**32
58
  np.random.seed(worker_seed)
59
  random.seed(worker_seed)
core/sample_ddim.py CHANGED
@@ -17,7 +17,7 @@ def ddim_sample(unet, noise_scheduler, shape, null_embedding=None, x_start=None,
17
  # x = torch.randn(shape, device=device) # start from pure noise
18
  noise_scheduler.set_timesteps(num_steps) # set DDIM timesteps
19
  embedding = embedding.expand(x.shape[0], -1, -1) if embedding != None else None
20
- if null_embedding == None: null_embedding = torch.load(f"{embedding_dir}/null_embedding.pt", map_location=device, weights_only=True).unsqueeze(0) # [1, 77, 1024]
21
  null_embedding = null_embedding.expand(x.shape[0], -1, -1)
22
  for t in tqdm(noise_scheduler.timesteps, desc=f"Sampling timesteps: ", colour=tqdm_colors[-1]):
23
  t_batch = torch.full((shape[0],), t, device=device, dtype=torch.long)
@@ -32,7 +32,7 @@ def ddim_sample(unet, noise_scheduler, shape, null_embedding=None, x_start=None,
32
  target_std = max(std_uncond, std_cond) # Target std is usually the max of the two individual predictions or just the conditional std.
33
  factor = target_std / std_cfg # Factor to bring the CFG output back into range
34
  model_out_rescaled = model_out * factor
35
- phi = 0.7 # Final blend (phi=0.7 is standard)
36
  model_out = phi * model_out_rescaled + (1 - phi) * model_out
37
  else: model_out = unet(x, t_batch, embedding)
38
  x = noise_scheduler.step(model_output=model_out, timestep=t, sample=x, eta=eta).prev_sample # DDIM step
 
17
  # x = torch.randn(shape, device=device) # start from pure noise
18
  noise_scheduler.set_timesteps(num_steps) # set DDIM timesteps
19
  embedding = embedding.expand(x.shape[0], -1, -1) if embedding != None else None
20
+ if null_embedding == None: null_embedding = torch.load(null_embedding_dir, map_location=device, weights_only=True).unsqueeze(0) # [1, 77, 1024]
21
  null_embedding = null_embedding.expand(x.shape[0], -1, -1)
22
  for t in tqdm(noise_scheduler.timesteps, desc=f"Sampling timesteps: ", colour=tqdm_colors[-1]):
23
  t_batch = torch.full((shape[0],), t, device=device, dtype=torch.long)
 
32
  target_std = max(std_uncond, std_cond) # Target std is usually the max of the two individual predictions or just the conditional std.
33
  factor = target_std / std_cfg # Factor to bring the CFG output back into range
34
  model_out_rescaled = model_out * factor
35
+ phi = 0.7 # Final blend
36
  model_out = phi * model_out_rescaled + (1 - phi) * model_out
37
  else: model_out = unet(x, t_batch, embedding)
38
  x = noise_scheduler.step(model_output=model_out, timestep=t, sample=x, eta=eta).prev_sample # DDIM step
core/text_embeddings.py CHANGED
@@ -11,10 +11,8 @@ def save_embeddings(model, tokenizer, device):
11
  last_img_name = ""; counter = 0
12
  for l in tqdm(f.readlines(), desc=f"Progress"):
13
  img_name, caption = l.strip().split(".jpg,")
14
- if caption[-1]==".":
15
- caption = caption.strip().rstrip('.')
16
- # caption = caption[:-2]
17
- caption = caption.lower()
18
  if len(last_img_name) == 0 or last_img_name != img_name: last_img_name = img_name; counter = 0
19
  else: counter += 1
20
  img_name += f"_{counter}"
@@ -24,7 +22,7 @@ def save_embeddings(model, tokenizer, device):
24
 
25
  def save_null_embedding(model, tokenizer, device):
26
  embedding = get_text_embedding(model, tokenizer, "", device)
27
- torch.save(embedding.cpu(), f"./{embedding_dir}/null_embedding.pt")
28
  print("Saved null embedding")
29
 
30
  def save_val_embeddings(model, tokenizer, device):
@@ -73,7 +71,7 @@ if __name__ == "__main__":
73
  save_val_embeddings(model, tokenizer, device)
74
 
75
  # Check if embeddings are equal
76
- # emb1 = torch.load("./backend/best_checkpoints/null_embedding/null_embedding.pt", map_location="cpu", weights_only=True)
77
  # emb2 = torch.load("./backend/core/embeddings_77_1024/null_embedding.pt", map_location="cpu", weights_only=True)
78
  # print("Max diff:", (emb1 - emb2).abs().max().item())
79
  # print("Are equal:", torch.allclose(emb2, emb1, atol=1e-5))
 
11
  last_img_name = ""; counter = 0
12
  for l in tqdm(f.readlines(), desc=f"Progress"):
13
  img_name, caption = l.strip().split(".jpg,")
14
+ caption = caption.strip()
15
+ if caption[-1]==".": caption = caption.rstrip('.')
 
 
16
  if len(last_img_name) == 0 or last_img_name != img_name: last_img_name = img_name; counter = 0
17
  else: counter += 1
18
  img_name += f"_{counter}"
 
22
 
23
  def save_null_embedding(model, tokenizer, device):
24
  embedding = get_text_embedding(model, tokenizer, "", device)
25
+ torch.save(embedding.cpu(), null_embedding_dir)
26
  print("Saved null embedding")
27
 
28
  def save_val_embeddings(model, tokenizer, device):
 
71
  save_val_embeddings(model, tokenizer, device)
72
 
73
  # Check if embeddings are equal
74
+ # emb1 = torch.load("./backend/core/null_embedding.pt", map_location="cpu", weights_only=True)
75
  # emb2 = torch.load("./backend/core/embeddings_77_1024/null_embedding.pt", map_location="cpu", weights_only=True)
76
  # print("Max diff:", (emb1 - emb2).abs().max().item())
77
  # print("Are equal:", torch.allclose(emb2, emb1, atol=1e-5))
core/train_unet.py CHANGED
@@ -16,7 +16,7 @@ from sample_ddim import gen_n_sampled_img, gen_val_sampled_img
16
  import numpy as np
17
  import random
18
 
19
- def train_unet_ddpm_simple(use_checkpoint=False):
20
  os.makedirs(unet_checkpoint_dir, exist_ok=True)
21
  device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
22
  print("DEVICE", device)
@@ -35,22 +35,19 @@ def train_unet_ddpm_simple(use_checkpoint=False):
35
  save_every = epoch_size * 8
36
  val_every = epoch_size * 4
37
  ema_warmup_steps = epoch_size * 20
38
- null_embedding = (torch.load(f"{embedding_dir}/null_embedding.pt", map_location=device, weights_only=True).unsqueeze(0))
39
  ema = EMA(unet, decay=0.999, warmup_steps=ema_warmup_steps)
40
  scaler = torch.amp.GradScaler("cuda", enabled=torch.cuda.is_available())
 
41
  global_step = loaded_step = 0
42
  total_mse = total_samples = 0
43
- checkpoint_path = "./backend/core/checkpoints/ldm/ema_step_586960_0.298282.pth"
44
- if use_checkpoint and os.path.exists(checkpoint_path):
45
- loaded_step = 586960
46
  print("Loaded checkpoint successfully")
47
  checkpoint = torch.load(checkpoint_path, map_location=device, weights_only=True)
48
  if "rng_state" in checkpoint: torch.set_rng_state(checkpoint["rng_state"].cpu())
49
  if torch.cuda.is_available() and "cuda_rng_state" in checkpoint:
50
  torch.cuda.set_rng_state(checkpoint["cuda_rng_state"].cpu())
51
- # torch.cuda.set_rng_state_all(checkpoint["cuda_rng_state"].cpu())
52
- # np.random.set_state(checkpoint["cuda_rng_state"].cpu())
53
- # random.setstate(checkpoint["cuda_rng_state"].cpu())
54
  unet.load_state_dict(checkpoint["unet"], strict=True)
55
  global_step = checkpoint["step"]
56
  optimizer.load_state_dict(checkpoint["optimizer"])
@@ -58,17 +55,19 @@ def train_unet_ddpm_simple(use_checkpoint=False):
58
  noise_scheduler = DDPMScheduler.from_pretrained("./checkpoints/noise_scheduler")
59
  ema.shadow = {k: v.to(device) for k, v in checkpoint["ema"].items()}
60
  ema.num_updates = checkpoint["ema_num_updates"]
 
61
  unet.train()
62
  while global_step < max_steps:
63
  for latent, embedding in tqdm(train_loader, desc="Train: ", colour=tqdm_colors[5]):
64
  if global_step >= max_steps: break
65
  latent, embedding = latent.to(device), embedding.to(device)
66
  assert null_embedding.shape[1:] == embedding.shape[1:], f"{null_embedding.shape}"
 
67
  batch_size = latent.size(0)
68
  t = torch.randint(0, noise_scheduler.config.num_train_timesteps, (batch_size,), device=device).long()
69
  noise = torch.randn_like(latent)
70
  noised_latent = noise_scheduler.add_noise(latent, noise, t)
71
- mask = torch.rand(batch_size, 1, 1, device=device) < 0.2
72
  embedding = torch.where(mask, null_embedding.expand(batch_size, -1, -1), embedding)
73
  optimizer.zero_grad(set_to_none=True)
74
  with torch.amp.autocast(device_type="cuda", enabled=torch.cuda.is_available()):
@@ -83,20 +82,14 @@ def train_unet_ddpm_simple(use_checkpoint=False):
83
  scaler.update()
84
  if scaler.get_scale() >= old_scale:
85
  ema.update(unet)
 
86
  global_step += 1
87
- # scaler.scale(loss).backward()
88
- # scaler.unscale_(optimizer)
89
- # torch.nn.utils.clip_grad_norm_(unet.parameters(), 1.0)
90
- # scaler.step(optimizer)
91
- # scaler.update()
92
- # # EMA update
93
- # ema.update(unet)
94
- # global_step += 1
95
- total_mse += loss.item() * batch_size
96
- total_samples += batch_size
97
- if global_step % log_every == 0 and (global_step != loaded_step):
98
- print(f" Step {global_step} | Loss: {total_mse / total_samples:.6f}")
99
- if global_step % save_every == 0 and (global_step != loaded_step):
100
  torch.save({
101
  "step": global_step,
102
  "unet": unet.state_dict(),
@@ -109,7 +102,7 @@ def train_unet_ddpm_simple(use_checkpoint=False):
109
  }, f"{unet_checkpoint_dir}/ema_step_{global_step}_{total_mse / total_samples:.6f}.pth")
110
  noise_scheduler.save_pretrained(f"./checkpoints/noise_scheduler")
111
  total_samples = total_mse = 0
112
- if global_step % val_every == 0 and (global_step != loaded_step):
113
  unet.eval()
114
  applied_ema = False
115
  with torch.no_grad():
@@ -126,4 +119,4 @@ def train_unet_ddpm_simple(use_checkpoint=False):
126
 
127
  if __name__ == "__main__":
128
  seed_everything(42)
129
- train_unet_ddpm_simple(use_checkpoint=False)
 
16
  import numpy as np
17
  import random
18
 
19
+ def train_unet_ddpm_simple(use_checkpoint=False, checkpoint_path=None):
20
  os.makedirs(unet_checkpoint_dir, exist_ok=True)
21
  device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
22
  print("DEVICE", device)
 
35
  save_every = epoch_size * 8
36
  val_every = epoch_size * 4
37
  ema_warmup_steps = epoch_size * 20
38
+ null_embedding = (torch.load(null_embedding_dir, map_location=device, weights_only=True).unsqueeze(0))
39
  ema = EMA(unet, decay=0.999, warmup_steps=ema_warmup_steps)
40
  scaler = torch.amp.GradScaler("cuda", enabled=torch.cuda.is_available())
41
+ # scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=max_steps, eta_min=1e-6)
42
  global_step = loaded_step = 0
43
  total_mse = total_samples = 0
44
+ if use_checkpoint and checkpoint_path and os.path.exists(checkpoint_path):
45
+ loaded_step = int(checkpoint_path.split("step_")[1].split("_")[0])
 
46
  print("Loaded checkpoint successfully")
47
  checkpoint = torch.load(checkpoint_path, map_location=device, weights_only=True)
48
  if "rng_state" in checkpoint: torch.set_rng_state(checkpoint["rng_state"].cpu())
49
  if torch.cuda.is_available() and "cuda_rng_state" in checkpoint:
50
  torch.cuda.set_rng_state(checkpoint["cuda_rng_state"].cpu())
 
 
 
51
  unet.load_state_dict(checkpoint["unet"], strict=True)
52
  global_step = checkpoint["step"]
53
  optimizer.load_state_dict(checkpoint["optimizer"])
 
55
  noise_scheduler = DDPMScheduler.from_pretrained("./checkpoints/noise_scheduler")
56
  ema.shadow = {k: v.to(device) for k, v in checkpoint["ema"].items()}
57
  ema.num_updates = checkpoint["ema_num_updates"]
58
+ # for _ in range(global_step): scheduler.step()
59
  unet.train()
60
  while global_step < max_steps:
61
  for latent, embedding in tqdm(train_loader, desc="Train: ", colour=tqdm_colors[5]):
62
  if global_step >= max_steps: break
63
  latent, embedding = latent.to(device), embedding.to(device)
64
  assert null_embedding.shape[1:] == embedding.shape[1:], f"{null_embedding.shape}"
65
+ assert embedding.dim() == 3, f"Expected 3D embedding (B, seq, dim), got {embedding.shape}"
66
  batch_size = latent.size(0)
67
  t = torch.randint(0, noise_scheduler.config.num_train_timesteps, (batch_size,), device=device).long()
68
  noise = torch.randn_like(latent)
69
  noised_latent = noise_scheduler.add_noise(latent, noise, t)
70
+ mask = torch.rand(batch_size, 1, 1, device=device) < 0.1
71
  embedding = torch.where(mask, null_embedding.expand(batch_size, -1, -1), embedding)
72
  optimizer.zero_grad(set_to_none=True)
73
  with torch.amp.autocast(device_type="cuda", enabled=torch.cuda.is_available()):
 
82
  scaler.update()
83
  if scaler.get_scale() >= old_scale:
84
  ema.update(unet)
85
+ # scheduler.step()
86
  global_step += 1
87
+ total_mse += loss.item() * batch_size
88
+ total_samples += batch_size
89
+ if global_step % log_every == 0 and (global_step != loaded_step):
90
+ # print(f" Step {global_step} | Loss: {total_mse / total_samples:.6f} | LR: {scheduler.get_last_lr()[0]:.2e}")
91
+ if total_samples > 0: print(f" Step {global_step} | Loss: {total_mse / total_samples:.6f}")
92
+ if global_step % save_every == 0 and (global_step != loaded_step):
 
 
 
 
 
 
 
93
  torch.save({
94
  "step": global_step,
95
  "unet": unet.state_dict(),
 
102
  }, f"{unet_checkpoint_dir}/ema_step_{global_step}_{total_mse / total_samples:.6f}.pth")
103
  noise_scheduler.save_pretrained(f"./checkpoints/noise_scheduler")
104
  total_samples = total_mse = 0
105
+ if global_step % val_every == 0 and (global_step != loaded_step):
106
  unet.eval()
107
  applied_ema = False
108
  with torch.no_grad():
 
119
 
120
  if __name__ == "__main__":
121
  seed_everything(42)
122
+ train_unet_ddpm_simple(use_checkpoint=True, checkpoint_path="./backend\core\checkpoints\ldm\ema_step_1163800_0.262618.pth")
core/train_vae.py CHANGED
@@ -1,5 +1,6 @@
1
  import sys, os
2
  sys.path.insert(0, os.path.dirname(__file__))
 
3
  import torch
4
  from torch import nn, optim
5
  import torch.nn.functional as F
@@ -14,12 +15,6 @@ from sample_vae import reconstruct
14
  import lpips
15
  from pytorch_msssim import SSIM
16
 
17
- # def vae_loss(recon_x, x, mu, logvar, beta_kld):
18
- # # loss = F.l1_loss(recon_x, x, reduction="sum") / recon_x.size(0)
19
- # loss = F.mse_loss(recon_x, x, reduction="sum") / recon_x.size(0)
20
- # kld = torch.mean(-0.5 * torch.sum(1 + logvar - mu.pow(2) - logvar.exp(), dim=1))
21
- # tvl = total_variance_loss(recon_x)
22
- # return loss + kld*beta_kld + tvl*vae_lambda_tvl, loss, kld, tvl
23
 
24
  # # Simple PatchGAN Discriminator
25
  # class PatchDiscriminator(nn.Module):
@@ -73,22 +68,25 @@ def total_variance_loss(x):
73
  tvl_w = torch.pow(x[:, :, :, 1:] - x[:, :, :, :-1], 2).sum() # TV-L2
74
  return (tvl_h + tvl_w) / x.size(0)
75
 
76
- def get_annealed_beta(epoch, warmup_epochs=100, max_beta=1.0): return max_beta * min(vae_beta_kld, epoch / warmup_epochs)
77
 
78
  # def vae_loss(recon_x, x, mu, logvar, lpips_fn, ssim_fn, beta_kld=1.0):
79
- def vae_loss(recon_x, x, mu, logvar, beta_kld=1.0):
80
  # mse = torch.mean((recon_x - x) ** 2)
81
  # kld_loss = torch.mean(-0.5 * (1 + logvar - mu.pow(2) - logvar.exp()))
82
  # b, c, h, w = x.shape
83
- # tvl = total_variance_loss(recon_x) if vae_lambda_tvl > 0 else 0.0
84
  mse = F.mse_loss(recon_x, x, reduction="sum") / recon_x.size(0)
85
  kld_loss = torch.sum(-0.5 * (1 + logvar - mu.pow(2) - logvar.exp())) / recon_x.size(0)
86
- # ssim_loss = 1 - ssim_fn(recon_x, x)
87
- # lpips_loss = lpips_fn(recon_x * 2 - 1, x * 2 - 1).mean()
88
- # total_loss = mse + (beta_kld * kld_loss) + (vae_lambda_tvl * tvl) + (vae_lpips_weight * lpips_loss) + ssim_loss * 0.1
89
- # return total_loss, mse, kld_loss, tvl, lpips_loss, ssim_loss
90
- total_loss = mse + (beta_kld * kld_loss)
91
- return total_loss, mse, kld_loss, 0, 0, 0
 
 
 
92
 
93
  def train_test_vae():
94
  device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
@@ -98,14 +96,14 @@ def train_test_vae():
98
  scheduler = ReduceLROnPlateau(optimizer, mode='min', patience=3, factor=0.5)
99
  os.makedirs(f"{vae_checkpoint_dir}", exist_ok=True)
100
  early_stopping_counter = 0
101
- # global_step = 0
102
- # warmup_steps = int(0.1 * len(train_loader) * vae_num_epochs)
103
  best_test_loss = float("inf")
104
  train_bce_loss = []
105
  train_kld_loss = []
106
- # lpips_fn = lpips.LPIPS(net='vgg').to(device)
107
- # lpips_fn.eval()
108
- # ssim_fn = SSIM(data_range=1, size_average=True, channel=3)
 
 
109
  scaler = torch.amp.GradScaler("cuda", enabled=torch.cuda.is_available())
110
  # discriminator = PatchDiscriminator().to(device)
111
  # optimizer_D = optim.AdamW(discriminator.parameters(), lr=vae_optim_lr, betas=(0.5, 0.999))
@@ -126,18 +124,19 @@ def train_test_vae():
126
  # loss, rec, kld, tvl, lpips_loss, gan_loss = vae_loss(recon_x, x, mu, logvar, discriminator, lpips_fn, beta_kld=get_annealed_beta(epoch)) # Update VAE (generator)
127
  # loss, rec, kld, tvl, lpips_loss = vae_loss(recon_x, x, mu, logvar, lpips_fn, beta_kld=min(1, global_step / warmup_steps)) # Update VAE (generator)
128
  # loss, rec, kld, tvl, lpips_loss, ssim_loss = vae_loss(recon_x, x, mu, logvar, lpips_fn, ssim_fn, beta_kld=get_annealed_beta(epoch)) # Update VAE (generator)
129
- loss, rec, kld, tvl, lpips_loss, ssim_loss = vae_loss(recon_x, x, mu, logvar, beta_kld=get_annealed_beta(epoch)) # Update VAE (generator)
130
  optimizer.zero_grad(set_to_none=True)
131
  scaler.scale(loss).backward()
 
 
132
  scaler.step(optimizer)
133
  scaler.update()
134
  train_loss += loss.item()
135
  total_rec += rec.item()
136
  total_kld += kld.item()
137
- # total_tvl += tvl
138
- # total_lpips += lpips_loss
139
- # total_lpips += lpips_loss.item()
140
- # total_ssim += ssim_loss.item()
141
  # total_gan += gan_loss.item()
142
  avg_train_loss = train_loss / len(train_loader)
143
  avg_rec = total_rec / len(train_loader)
@@ -157,7 +156,7 @@ def train_test_vae():
157
  recon_x, mu, logvar = vae(x)
158
  # loss, *_ = vae_loss(recon_x, x, mu, logvar, discriminator, lpips_fn, beta_kld=1.0)
159
  # loss, *_ = vae_loss(recon_x, x, mu, logvar, lpips_fn, ssim_fn, beta_kld=vae_beta_kld)
160
- loss, *_ = vae_loss(recon_x, x, mu, logvar, beta_kld=vae_beta_kld)
161
  test_loss += loss.item()
162
  avg_test_loss = test_loss / len(test_loader)
163
  scheduler.step(avg_test_loss)
@@ -187,7 +186,6 @@ def plot_recon_vs_kld(train_bce_loss, train_kld_loss):
187
  plt.figure(figsize=(10, 6))
188
  plt.plot(epochs, train_bce_loss, label='Reconstruction Loss (BCE)', color='blue', linewidth=2)
189
  plt.plot(epochs, train_kld_loss, label='KL Divergence', color='red', linewidth=2)
190
- # plt.plot(epochs, train_tvl_loss, label='TVL', color='yellow', linewidth=2)
191
  plt.xlabel('Epoch')
192
  plt.ylabel('Loss')
193
  plt.title('Reconstruction Loss vs KL Divergence over Epochs')
 
1
  import sys, os
2
  sys.path.insert(0, os.path.dirname(__file__))
3
+ os.environ["KMP_DUPLICATE_LIB_OK"] = "TRUE"
4
  import torch
5
  from torch import nn, optim
6
  import torch.nn.functional as F
 
15
  import lpips
16
  from pytorch_msssim import SSIM
17
 
 
 
 
 
 
 
18
 
19
  # # Simple PatchGAN Discriminator
20
  # class PatchDiscriminator(nn.Module):
 
68
  tvl_w = torch.pow(x[:, :, :, 1:] - x[:, :, :, :-1], 2).sum() # TV-L2
69
  return (tvl_h + tvl_w) / x.size(0)
70
 
71
+ def get_annealed_beta(epoch, warmup_epochs=100, max_beta=1.0): return vae_beta_kld * min(max_beta, epoch / warmup_epochs)
72
 
73
  # def vae_loss(recon_x, x, mu, logvar, lpips_fn, ssim_fn, beta_kld=1.0):
74
+ def vae_loss(recon_x, x, mu, logvar, beta_kld=1.0, lpips_fn=None, ssim_fn=None):
75
  # mse = torch.mean((recon_x - x) ** 2)
76
  # kld_loss = torch.mean(-0.5 * (1 + logvar - mu.pow(2) - logvar.exp()))
77
  # b, c, h, w = x.shape
78
+ tvl = total_variance_loss(recon_x) if vae_lambda_tvl > 0 else 0.0
79
  mse = F.mse_loss(recon_x, x, reduction="sum") / recon_x.size(0)
80
  kld_loss = torch.sum(-0.5 * (1 + logvar - mu.pow(2) - logvar.exp())) / recon_x.size(0)
81
+ # mse = F.mse_loss(recon_x, x, reduction="mean")
82
+ # kld_loss = torch.mean(-0.5 * (1 + logvar - mu.pow(2) - logvar.exp())) # divides by latent_ch×H×W×B
83
+ with torch.amp.autocast("cuda", enabled=False):
84
+ ssim_loss = 1 - ssim_fn(recon_x.float(), x.float())
85
+ lpips_loss = lpips_fn(recon_x.float() * 2 - 1, x.float() * 2 - 1).mean()
86
+ total_loss = mse + (beta_kld * kld_loss) + (vae_lambda_tvl * tvl) + (vae_lpips_weight * lpips_loss) + ssim_loss * 0.1
87
+ return total_loss, mse, kld_loss, tvl, lpips_loss, ssim_loss
88
+ # total_loss = mse + (beta_kld * kld_loss)
89
+ # return total_loss, mse, kld_loss, 0, 0, 0
90
 
91
  def train_test_vae():
92
  device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
 
96
  scheduler = ReduceLROnPlateau(optimizer, mode='min', patience=3, factor=0.5)
97
  os.makedirs(f"{vae_checkpoint_dir}", exist_ok=True)
98
  early_stopping_counter = 0
 
 
99
  best_test_loss = float("inf")
100
  train_bce_loss = []
101
  train_kld_loss = []
102
+ lpips_fn = lpips.LPIPS(net='vgg').to(device)
103
+ lpips_fn.eval()
104
+ for param in lpips_fn.parameters():
105
+ param.requires_grad = False
106
+ ssim_fn = SSIM(data_range=1, size_average=True, channel=3)
107
  scaler = torch.amp.GradScaler("cuda", enabled=torch.cuda.is_available())
108
  # discriminator = PatchDiscriminator().to(device)
109
  # optimizer_D = optim.AdamW(discriminator.parameters(), lr=vae_optim_lr, betas=(0.5, 0.999))
 
124
  # loss, rec, kld, tvl, lpips_loss, gan_loss = vae_loss(recon_x, x, mu, logvar, discriminator, lpips_fn, beta_kld=get_annealed_beta(epoch)) # Update VAE (generator)
125
  # loss, rec, kld, tvl, lpips_loss = vae_loss(recon_x, x, mu, logvar, lpips_fn, beta_kld=min(1, global_step / warmup_steps)) # Update VAE (generator)
126
  # loss, rec, kld, tvl, lpips_loss, ssim_loss = vae_loss(recon_x, x, mu, logvar, lpips_fn, ssim_fn, beta_kld=get_annealed_beta(epoch)) # Update VAE (generator)
127
+ loss, rec, kld, tvl, lpips_loss, ssim_loss = vae_loss(recon_x, x, mu, logvar, beta_kld=get_annealed_beta(epoch), lpips_fn=lpips_fn, ssim_fn=ssim_fn) # Update VAE (generator)
128
  optimizer.zero_grad(set_to_none=True)
129
  scaler.scale(loss).backward()
130
+ scaler.unscale_(optimizer)
131
+ torch.nn.utils.clip_grad_norm_(vae.parameters(), max_norm=1.0)
132
  scaler.step(optimizer)
133
  scaler.update()
134
  train_loss += loss.item()
135
  total_rec += rec.item()
136
  total_kld += kld.item()
137
+ total_tvl += tvl.item() if vae_lambda_tvl > 0 else 0.0
138
+ total_lpips += lpips_loss.item()
139
+ total_ssim += ssim_loss.item()
 
140
  # total_gan += gan_loss.item()
141
  avg_train_loss = train_loss / len(train_loader)
142
  avg_rec = total_rec / len(train_loader)
 
156
  recon_x, mu, logvar = vae(x)
157
  # loss, *_ = vae_loss(recon_x, x, mu, logvar, discriminator, lpips_fn, beta_kld=1.0)
158
  # loss, *_ = vae_loss(recon_x, x, mu, logvar, lpips_fn, ssim_fn, beta_kld=vae_beta_kld)
159
+ loss, *_ = vae_loss(recon_x, x, mu, logvar, beta_kld=vae_beta_kld, lpips_fn=lpips_fn, ssim_fn=ssim_fn)
160
  test_loss += loss.item()
161
  avg_test_loss = test_loss / len(test_loader)
162
  scheduler.step(avg_test_loss)
 
186
  plt.figure(figsize=(10, 6))
187
  plt.plot(epochs, train_bce_loss, label='Reconstruction Loss (BCE)', color='blue', linewidth=2)
188
  plt.plot(epochs, train_kld_loss, label='KL Divergence', color='red', linewidth=2)
 
189
  plt.xlabel('Epoch')
190
  plt.ylabel('Loss')
191
  plt.title('Reconstruction Loss vs KL Divergence over Epochs')
core/unet.py CHANGED
@@ -12,8 +12,8 @@ def get_timestep_embedding(timesteps: torch.Tensor, dim: int = 1024):
12
  """Sinusoidal position embeddings"""
13
  assert timesteps.dim() == 1
14
  half_dim = dim // 2
15
- freqs = torch.exp(-math.log(10000) * torch.arange(0, half_dim, dtype=torch.float32, device=timesteps.device) / (half_dim - 1))
16
- # freqs = torch.exp(-math.log(10000) * torch.arange(0, half_dim, dtype=torch.float32, device=timesteps.device) / (half_dim))
17
  args = timesteps.float().unsqueeze(1) * freqs.unsqueeze(0)
18
  emb = torch.cat([torch.sin(args), torch.cos(args)], dim=-1)
19
  if dim % 2: emb = F.pad(emb, (0, 1)) # zero‑pad if dim is odd
@@ -32,7 +32,7 @@ class ResBlock(nn.Module):
32
  super().__init__()
33
  self.convgnact1 = ConvGNAct(in_ch, out_ch)
34
  self.norm = nn.GroupNorm(unet_group_size, out_ch)
35
- self.act = nn.SiLU(inplace=True)
36
  self.out_conv = nn.Conv2d(out_ch, out_ch, 3, padding=1)
37
  self.time_emb_proj = nn.Sequential(
38
  nn.SiLU(),
@@ -60,26 +60,17 @@ class DownBlock(nn.Module):
60
  super().__init__()
61
  self.resblock1 = ResBlock(in_ch, out_ch, time_emb_dim)
62
  self.resblock2 = ResBlock(out_ch, out_ch, time_emb_dim)
63
- self.resblock3 = ResBlock(out_ch, out_ch, time_emb_dim)
64
- self.resblock4 = ResBlock(out_ch, out_ch, time_emb_dim)
65
- # self.resblock5 = ResBlock(out_ch, out_ch, time_emb_dim)
66
- # self.resblock6 = ResBlock(out_ch, out_ch, time_emb_dim)
67
- # self.resblock7 = ResBlock(out_ch, out_ch, time_emb_dim)
68
- # self.resblock8 = ResBlock(out_ch, out_ch, time_emb_dim)
69
- self.attn = SelfCrossAttn(out_ch, heads=heads, text_emb_dim=text_emb_dim, cross=cross)
70
  # self.down = ConvGNAct(out_ch, out_ch, 3, 2, 1)
71
  self.down = nn.Conv2d(out_ch, out_ch, 3, 2, 1)
72
  self.use_attn = use_attn
73
  def forward(self, x, time_emb, text):
74
  x = self.resblock1(x, time_emb)
75
  x = self.resblock2(x, time_emb)
76
- x = self.resblock3(x, time_emb)
77
- x = self.resblock4(x, time_emb)
78
- # x = self.resblock5(x, time_emb)
79
- # x = self.resblock6(x, time_emb)
80
- # x = self.resblock7(x, time_emb)
81
- # x = self.resblock8(x, time_emb)
82
- if self.use_attn: x = self.attn(x, text)
83
  x_down = self.down(x)
84
  return x_down, x
85
  class UpBlock(nn.Module):
@@ -88,13 +79,8 @@ class UpBlock(nn.Module):
88
  # self.up = nn.Conv2d(in_ch, in_ch, kernel_size=3, padding=1)
89
  self.resblock1 = ResBlock(in_ch * 2, out_ch, time_emb_dim)
90
  self.resblock2 = ResBlock(out_ch, out_ch, time_emb_dim)
91
- self.resblock3 = ResBlock(out_ch, out_ch, time_emb_dim)
92
- self.resblock4 = ResBlock(out_ch, out_ch, time_emb_dim)
93
- # self.resblock5 = ResBlock(out_ch, out_ch, time_emb_dim)
94
- # self.resblock6 = ResBlock(out_ch, out_ch, time_emb_dim)
95
- # self.resblock7 = ResBlock(out_ch, out_ch, time_emb_dim)
96
- # self.resblock8 = ResBlock(out_ch, out_ch, time_emb_dim)
97
- self.attn = SelfCrossAttn(out_ch, heads=heads, text_emb_dim=text_emb_dim, cross=cross)
98
  # self.out_proj = ConvGNAct(out_ch, out_ch, k=1, s=1, p=0)
99
  self.use_attn = use_attn
100
  def forward(self, x, skip, time_emb, text):
@@ -104,13 +90,9 @@ class UpBlock(nn.Module):
104
  x = torch.cat([x, skip], dim=1)
105
  x = self.resblock1(x, time_emb)
106
  x = self.resblock2(x, time_emb)
107
- x = self.resblock3(x, time_emb)
108
- x = self.resblock4(x, time_emb)
109
- # x = self.resblock5(x, time_emb)
110
- # x = self.resblock6(x, time_emb)
111
- # x = self.resblock7(x, time_emb)
112
- # x = self.resblock8(x, time_emb)
113
- if self.use_attn: x = self.attn(x, text)
114
  # x = self.out_proj(x)
115
  return x
116
  class Unet(nn.Module):
@@ -125,13 +107,14 @@ class Unet(nn.Module):
125
  # Downsample
126
  self.down1 = DownBlock(in_ch=base_ch, out_ch=base_ch * 2, time_emb_dim=time_emb_dim, text_emb_dim=text_emb_dim, cross=True, use_attn=False, heads=4)
127
  self.down2 = DownBlock(in_ch=base_ch * 2, out_ch=base_ch * 4, time_emb_dim=time_emb_dim, text_emb_dim=text_emb_dim, cross=True, use_attn=True, heads=8)
128
- self.down3 = DownBlock(in_ch=base_ch * 4, out_ch=base_ch * 8, time_emb_dim=time_emb_dim, text_emb_dim=text_emb_dim, cross=True, use_attn=True, heads=8)
129
  # Mid (bottleneck)
130
  self.mid_res1 = ResBlock(in_ch=base_ch * 8, out_ch=base_ch * 8, time_emb_dim=time_emb_dim)
131
- self.mid_attn = SelfCrossAttn(base_ch * 8, heads=16, text_emb_dim=1024, cross=True)
 
132
  self.mid_res2 = ResBlock(in_ch=base_ch * 8, out_ch=base_ch * 8, time_emb_dim=time_emb_dim)
133
  # Upsample
134
- self.up3 = UpBlock(in_ch=base_ch * 8, out_ch=base_ch * 4, time_emb_dim=time_emb_dim, text_emb_dim=text_emb_dim, cross=True, use_attn=True, heads=8)
135
  self.up2 = UpBlock(in_ch=base_ch * 4, out_ch=base_ch * 2, time_emb_dim=time_emb_dim, text_emb_dim=text_emb_dim, cross=True, use_attn=True, heads=8)
136
  self.up1 = UpBlock(in_ch=base_ch * 2, out_ch=base_ch, time_emb_dim=time_emb_dim, text_emb_dim=text_emb_dim, cross=True, use_attn=False, heads=4)
137
  self.out_norm = nn.GroupNorm(unet_group_size, base_ch)
@@ -144,7 +127,8 @@ class Unet(nn.Module):
144
  h, skip2 = self.down2(h, time_emb, text_emb) # (B, C x 4, H / 4, W / 4) (B, C x 4, H / 2, W / 2) # print("DOWN 2", h.shape, skip2.shape)
145
  h, skip3 = self.down3(h, time_emb, text_emb) # (B, C x 8, H / 8, W / 8) (B, C x 4, H / 4, W / 4) # print("DOWN 3", h.shape, skip3.shape)
146
  h = self.mid_res1(h, time_emb) # (B, C x 8, H / 8, W / 8) # print("MID RES", h.shape)
147
- h = self.mid_attn(h, text_emb) # (B, C x 8, H / 8, W / 8) # print("MID ATTN", h.shape)
 
148
  h = self.mid_res2(h, time_emb) # (B, C x 8, H / 8, W / 8) # print("MID RES", h.shape)
149
  h = self.up3(h, skip3, time_emb, text_emb) # (B, C x 4, H / 4, W / 4) # print("UP 2", h.shape)
150
  h = self.up2(h, skip2, time_emb, text_emb) # (B, C x 2, H / 2, W / 2) # print("UP 2", h.shape)
 
12
  """Sinusoidal position embeddings"""
13
  assert timesteps.dim() == 1
14
  half_dim = dim // 2
15
+ # freqs = torch.exp(-math.log(10000) * torch.arange(0, half_dim, dtype=torch.float32, device=timesteps.device) / (half_dim - 1))
16
+ freqs = torch.exp(-math.log(10000) * torch.arange(0, half_dim, dtype=torch.float32, device=timesteps.device) / (half_dim))
17
  args = timesteps.float().unsqueeze(1) * freqs.unsqueeze(0)
18
  emb = torch.cat([torch.sin(args), torch.cos(args)], dim=-1)
19
  if dim % 2: emb = F.pad(emb, (0, 1)) # zero‑pad if dim is odd
 
32
  super().__init__()
33
  self.convgnact1 = ConvGNAct(in_ch, out_ch)
34
  self.norm = nn.GroupNorm(unet_group_size, out_ch)
35
+ self.act = nn.SiLU()
36
  self.out_conv = nn.Conv2d(out_ch, out_ch, 3, padding=1)
37
  self.time_emb_proj = nn.Sequential(
38
  nn.SiLU(),
 
60
  super().__init__()
61
  self.resblock1 = ResBlock(in_ch, out_ch, time_emb_dim)
62
  self.resblock2 = ResBlock(out_ch, out_ch, time_emb_dim)
63
+ self.self_attn = SelfCrossAttn(out_ch, heads=heads, text_emb_dim=text_emb_dim, cross=False)
64
+ self.cross_attn = SelfCrossAttn(out_ch, heads=heads, text_emb_dim=text_emb_dim, cross=True)
 
 
 
 
 
65
  # self.down = ConvGNAct(out_ch, out_ch, 3, 2, 1)
66
  self.down = nn.Conv2d(out_ch, out_ch, 3, 2, 1)
67
  self.use_attn = use_attn
68
  def forward(self, x, time_emb, text):
69
  x = self.resblock1(x, time_emb)
70
  x = self.resblock2(x, time_emb)
71
+ if self.use_attn:
72
+ x = self.self_attn(x, text)
73
+ x = self.cross_attn(x, text)
 
 
 
 
74
  x_down = self.down(x)
75
  return x_down, x
76
  class UpBlock(nn.Module):
 
79
  # self.up = nn.Conv2d(in_ch, in_ch, kernel_size=3, padding=1)
80
  self.resblock1 = ResBlock(in_ch * 2, out_ch, time_emb_dim)
81
  self.resblock2 = ResBlock(out_ch, out_ch, time_emb_dim)
82
+ self.self_attn = SelfCrossAttn(out_ch, heads=heads, text_emb_dim=text_emb_dim, cross=False)
83
+ self.cross_attn = SelfCrossAttn(out_ch, heads=heads, text_emb_dim=text_emb_dim, cross=True)
 
 
 
 
 
84
  # self.out_proj = ConvGNAct(out_ch, out_ch, k=1, s=1, p=0)
85
  self.use_attn = use_attn
86
  def forward(self, x, skip, time_emb, text):
 
90
  x = torch.cat([x, skip], dim=1)
91
  x = self.resblock1(x, time_emb)
92
  x = self.resblock2(x, time_emb)
93
+ if self.use_attn:
94
+ x = self.self_attn(x, text)
95
+ x = self.cross_attn(x, text)
 
 
 
 
96
  # x = self.out_proj(x)
97
  return x
98
  class Unet(nn.Module):
 
107
  # Downsample
108
  self.down1 = DownBlock(in_ch=base_ch, out_ch=base_ch * 2, time_emb_dim=time_emb_dim, text_emb_dim=text_emb_dim, cross=True, use_attn=False, heads=4)
109
  self.down2 = DownBlock(in_ch=base_ch * 2, out_ch=base_ch * 4, time_emb_dim=time_emb_dim, text_emb_dim=text_emb_dim, cross=True, use_attn=True, heads=8)
110
+ self.down3 = DownBlock(in_ch=base_ch * 4, out_ch=base_ch * 8, time_emb_dim=time_emb_dim, text_emb_dim=text_emb_dim, cross=True, use_attn=True, heads=16)
111
  # Mid (bottleneck)
112
  self.mid_res1 = ResBlock(in_ch=base_ch * 8, out_ch=base_ch * 8, time_emb_dim=time_emb_dim)
113
+ self.mid_self_attn = SelfCrossAttn(base_ch * 8, heads=32, text_emb_dim=1024, cross=False)
114
+ self.mid_cross_attn = SelfCrossAttn(base_ch * 8, heads=32, text_emb_dim=1024, cross=True)
115
  self.mid_res2 = ResBlock(in_ch=base_ch * 8, out_ch=base_ch * 8, time_emb_dim=time_emb_dim)
116
  # Upsample
117
+ self.up3 = UpBlock(in_ch=base_ch * 8, out_ch=base_ch * 4, time_emb_dim=time_emb_dim, text_emb_dim=text_emb_dim, cross=True, use_attn=True, heads=16)
118
  self.up2 = UpBlock(in_ch=base_ch * 4, out_ch=base_ch * 2, time_emb_dim=time_emb_dim, text_emb_dim=text_emb_dim, cross=True, use_attn=True, heads=8)
119
  self.up1 = UpBlock(in_ch=base_ch * 2, out_ch=base_ch, time_emb_dim=time_emb_dim, text_emb_dim=text_emb_dim, cross=True, use_attn=False, heads=4)
120
  self.out_norm = nn.GroupNorm(unet_group_size, base_ch)
 
127
  h, skip2 = self.down2(h, time_emb, text_emb) # (B, C x 4, H / 4, W / 4) (B, C x 4, H / 2, W / 2) # print("DOWN 2", h.shape, skip2.shape)
128
  h, skip3 = self.down3(h, time_emb, text_emb) # (B, C x 8, H / 8, W / 8) (B, C x 4, H / 4, W / 4) # print("DOWN 3", h.shape, skip3.shape)
129
  h = self.mid_res1(h, time_emb) # (B, C x 8, H / 8, W / 8) # print("MID RES", h.shape)
130
+ h = self.mid_self_attn(h, text_emb) # (B, C x 8, H / 8, W / 8) # print("MID ATTN", h.shape)
131
+ h = self.mid_cross_attn(h, text_emb) # (B, C x 8, H / 8, W / 8) # print("MID ATTN", h.shape)
132
  h = self.mid_res2(h, time_emb) # (B, C x 8, H / 8, W / 8) # print("MID RES", h.shape)
133
  h = self.up3(h, skip3, time_emb, text_emb) # (B, C x 4, H / 4, W / 4) # print("UP 2", h.shape)
134
  h = self.up2(h, skip2, time_emb, text_emb) # (B, C x 2, H / 2, W / 2) # print("UP 2", h.shape)
core/vae.py CHANGED
@@ -6,32 +6,50 @@ import torch.nn.functional as F
6
  from config import *
7
  from attention import SelfCrossAttn
8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
9
  class VAE(nn.Module):
10
  def __init__(self):
11
  super().__init__()
12
  # Encoder
13
  self.encoder_conv = nn.Sequential(
14
- nn.Conv2d(3, 32, 3, 1, padding=1, bias=False), # (B, 4, 128, 128)
15
  nn.GroupNorm(vae_group_size, 32), nn.SiLU(inplace=True),
16
- nn.Conv2d(32, 128, 3, 2, padding=1, bias=False), # (B, 32, 64, 64)
17
  nn.GroupNorm(vae_group_size, 128), nn.SiLU(inplace=True),
18
- nn.Conv2d(128, 256, 3, 2, padding=1, bias=False), # (B, 128, 32, 32)
19
  nn.GroupNorm(vae_group_size, 256), nn.SiLU(inplace=True),
 
 
20
  # nn.Conv2d(256, 512, 3, 2, padding=1, bias=False), # (B, 128, 16, 16)
21
  # nn.GroupNorm(vae_group_size, 512), nn.SiLU(inplace=True),
22
- SelfCrossAttn(256, heads=8, cross=False),
 
 
23
  )
24
- # 1×1×1 convs give channel‑wise μ and log σ², shape = (B, latent_channels, 4, 4, 4)
25
- self.to_latent = nn.Conv2d(256, 2 * vae_latent_channels, kernel_size=1)
26
- self.conv_mu = nn.Conv2d(vae_latent_channels, vae_latent_channels, kernel_size=1)
27
- self.conv_logvar = nn.Conv2d(vae_latent_channels, vae_latent_channels, kernel_size=1)
28
- self.from_latent = nn.Conv2d(vae_latent_channels, 256, kernel_size=1)
29
  # Decoder
30
  self.decoder_conv = nn.Sequential(
31
- SelfCrossAttn(256, heads=8, cross=False),
32
- # nn.Upsample(scale_factor=2, mode="nearest"),
33
- # nn.Conv2d(512, 256, 3, padding=1, bias=False),
34
- # nn.GroupNorm(vae_group_size, 256), nn.SiLU(inplace=True),
 
 
35
  nn.Upsample(scale_factor=2, mode="nearest"),
36
  nn.Conv2d(256, 128, 3, padding=1, bias=False),
37
  nn.GroupNorm(vae_group_size, 128), nn.SiLU(inplace=True),
@@ -47,13 +65,12 @@ class VAE(nn.Module):
47
  std = torch.exp(0.5 * logvar)
48
  eps = torch.randn_like(std)
49
  return mu + eps * std
 
50
  def forward(self, x):
51
  h = self.encoder_conv(x) # (B, C, D, H)
52
  h = self.to_latent(h)
53
- # mu, logvar = self.conv_mu(h), self.conv_logvar(h)
54
  mu, logvar = torch.chunk(h, 2, dim=1)
55
  z = self.reparameterize(mu, logvar) # Latent (B, C, D, H)
56
- # z = z * 0.18215
57
  h = self.from_latent(z)
58
  recon = self.decoder_conv(h)
59
  return recon, mu, logvar
@@ -61,7 +78,6 @@ class VAE(nn.Module):
61
  def encode_img_to_latent(self, x):
62
  h = self.encoder_conv(x) # (B, C, D, H)
63
  h = self.to_latent(h)
64
- # mu, logvar = self.conv_mu(h), self.conv_logvar(h)
65
  mu, logvar = torch.chunk(h, 2, dim=1)
66
  z = self.reparameterize(mu, logvar)
67
  return z
 
6
  from config import *
7
  from attention import SelfCrossAttn
8
 
9
+ class VAEResBlock(nn.Module):
10
+ def __init__(self, in_channels, out_channels=None):
11
+ out_channels = out_channels or in_channels
12
+ super().__init__()
13
+ self.block = nn.Sequential(
14
+ nn.GroupNorm(vae_group_size, in_channels), nn.SiLU(inplace=True),
15
+ nn.Conv2d(in_channels, out_channels, 3, padding=1, bias=False),
16
+ nn.GroupNorm(vae_group_size, out_channels), nn.SiLU(inplace=True),
17
+ nn.Conv2d(out_channels, out_channels, 3, padding=1, bias=False),
18
+ )
19
+ self.skip = nn.Conv2d(in_channels, out_channels, 1, bias=False) if in_channels != out_channels else nn.Identity()
20
+ def forward(self, x):
21
+ return self.block(x) + self.skip(x)
22
+
23
  class VAE(nn.Module):
24
  def __init__(self):
25
  super().__init__()
26
  # Encoder
27
  self.encoder_conv = nn.Sequential(
28
+ nn.Conv2d(3, 32, 3, 1, padding=1, bias=False), # (B, 4, 256, 256)
29
  nn.GroupNorm(vae_group_size, 32), nn.SiLU(inplace=True),
30
+ nn.Conv2d(32, 128, 3, 2, padding=1, bias=False), # (B, 32, 128, 128)
31
  nn.GroupNorm(vae_group_size, 128), nn.SiLU(inplace=True),
32
+ nn.Conv2d(128, 256, 3, 2, padding=1, bias=False), # (B, 128, 64, 64)
33
  nn.GroupNorm(vae_group_size, 256), nn.SiLU(inplace=True),
34
+ nn.Conv2d(256, 512, 3, 2, padding=1, bias=False), # (B, 128, 32, 32)
35
+ nn.GroupNorm(vae_group_size, 512), nn.SiLU(inplace=True),
36
  # nn.Conv2d(256, 512, 3, 2, padding=1, bias=False), # (B, 128, 16, 16)
37
  # nn.GroupNorm(vae_group_size, 512), nn.SiLU(inplace=True),
38
+ # SelfCrossAttn(512, heads=8, cross=False),
39
+ VAEResBlock(512), SelfCrossAttn(512, heads=8, cross=False), VAEResBlock(512),
40
+ nn.GroupNorm(vae_group_size, 512), nn.SiLU(inplace=True),
41
  )
42
+ # Channel‑wise μ and log σ², shape = (B, latent_channels, 4, 4, 4)
43
+ self.to_latent = nn.Conv2d(512, 2 * vae_latent_channels, kernel_size=1)
44
+ self.from_latent = nn.Conv2d(vae_latent_channels, 512, kernel_size=1)
 
 
45
  # Decoder
46
  self.decoder_conv = nn.Sequential(
47
+ VAEResBlock(512), SelfCrossAttn(512, heads=8, cross=False), VAEResBlock(512),
48
+ nn.GroupNorm(vae_group_size, 512), nn.SiLU(inplace=True),
49
+ # SelfCrossAttn(512, heads=8, cross=False),
50
+ nn.Upsample(scale_factor=2, mode="nearest"),
51
+ nn.Conv2d(512, 256, 3, padding=1, bias=False),
52
+ nn.GroupNorm(vae_group_size, 256), nn.SiLU(inplace=True),
53
  nn.Upsample(scale_factor=2, mode="nearest"),
54
  nn.Conv2d(256, 128, 3, padding=1, bias=False),
55
  nn.GroupNorm(vae_group_size, 128), nn.SiLU(inplace=True),
 
65
  std = torch.exp(0.5 * logvar)
66
  eps = torch.randn_like(std)
67
  return mu + eps * std
68
+
69
  def forward(self, x):
70
  h = self.encoder_conv(x) # (B, C, D, H)
71
  h = self.to_latent(h)
 
72
  mu, logvar = torch.chunk(h, 2, dim=1)
73
  z = self.reparameterize(mu, logvar) # Latent (B, C, D, H)
 
74
  h = self.from_latent(z)
75
  recon = self.decoder_conv(h)
76
  return recon, mu, logvar
 
78
  def encode_img_to_latent(self, x):
79
  h = self.encoder_conv(x) # (B, C, D, H)
80
  h = self.to_latent(h)
 
81
  mu, logvar = torch.chunk(h, 2, dim=1)
82
  z = self.reparameterize(mu, logvar)
83
  return z