Spaces:
Sleeping
Sleeping
Updated: VAE, UNet, config, text embeddings, model and main
Browse files- app/main.py +4 -77
- app/model.py +3 -5
- core/config.py +11 -7
- core/dataloader.py +10 -9
- core/sample_ddim.py +2 -2
- core/text_embeddings.py +4 -6
- core/train_unet.py +17 -24
- core/train_vae.py +25 -27
- core/unet.py +19 -35
- core/vae.py +32 -16
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(
|
| 40 |
-
guidance_scale: float = Field(
|
| 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=
|
| 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 =
|
| 4 |
image_dir = f"{data_dir}/Images"
|
| 5 |
resized_img_dir = f"{data_dir}/ResizedImages_{image_res}"
|
| 6 |
-
vae_batch_size =
|
| 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 =
|
| 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}
|
|
|
|
| 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 =
|
| 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,
|
| 30 |
-
test_loader = DataLoader(test_set, batch_size=vae_batch_size, shuffle=False,
|
| 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 =
|
|
|
|
| 49 |
dataset = LatentEmbeddingsDataset(file_names)
|
| 50 |
-
|
| 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(
|
| 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
|
| 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 |
-
|
| 15 |
-
|
| 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(),
|
| 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/
|
| 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(
|
| 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
|
| 44 |
-
|
| 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.
|
| 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 |
-
|
| 88 |
-
|
| 89 |
-
|
| 90 |
-
|
| 91 |
-
|
| 92 |
-
|
| 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
|
| 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=
|
|
|
|
| 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
|
| 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 |
-
|
| 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 |
-
#
|
| 87 |
-
#
|
| 88 |
-
|
| 89 |
-
|
| 90 |
-
|
| 91 |
-
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
| 107 |
-
|
| 108 |
-
|
|
|
|
|
|
|
| 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 |
-
|
| 138 |
-
|
| 139 |
-
|
| 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 |
-
|
| 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(
|
| 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.
|
| 64 |
-
self.
|
| 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 |
-
|
| 77 |
-
|
| 78 |
-
|
| 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.
|
| 92 |
-
self.
|
| 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 |
-
|
| 108 |
-
|
| 109 |
-
|
| 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=
|
| 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.
|
|
|
|
| 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=
|
| 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.
|
|
|
|
| 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,
|
| 15 |
nn.GroupNorm(vae_group_size, 32), nn.SiLU(inplace=True),
|
| 16 |
-
nn.Conv2d(32, 128, 3, 2, padding=1, bias=False), # (B, 32,
|
| 17 |
nn.GroupNorm(vae_group_size, 128), nn.SiLU(inplace=True),
|
| 18 |
-
nn.Conv2d(128, 256, 3, 2, padding=1, bias=False), # (B, 128,
|
| 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(
|
|
|
|
|
|
|
| 23 |
)
|
| 24 |
-
#
|
| 25 |
-
self.to_latent = nn.Conv2d(
|
| 26 |
-
self.
|
| 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(
|
| 32 |
-
|
| 33 |
-
#
|
| 34 |
-
|
|
|
|
|
|
|
| 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
|