Skip to content

Commit e86eaf9

Browse files
authored
fix: DDIM clamp opt-in, negative prompt wiring, EMA bug, latent_dir, memory_format flag (#2)
* chore: add MIT License, CITATION.cff, .env.example, expand .gitignore * fix: DDIM clamp opt-in, negative prompt wiring, EMA arg, latent_dir default, memory_format flag - DDIMScheduler: make pred_x0 clamp opt-in via clamp_pred_x0 param (default False) - inference.py: wire --negative prompt through to uncond embedding - SD_ImageGen.py: pass model.unet to ema.apply_shadow() - SD_Train.py/train.py/SD_Train_v2.py: fix latent_dir default (remove double nesting) - SD_Train.py/train.py: add --memory_format CLI flag to gate channels_last (use 'contiguous' for AMD/Apple Silicon)
1 parent 1f5459b commit e86eaf9

7 files changed

Lines changed: 67 additions & 44 deletions

File tree

src/SD_ImageGen.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -82,7 +82,7 @@ def load_model(checkpoint_path: str, device: torch.device) -> StableDiffusionMod
8282
ema = EMA(model.unet, decay=0.9999)
8383
ema.load_state_dict(ckpt["ema_state_dict"])
8484
# Apply EMA shadow weights to the model
85-
ema.apply_shadow()
85+
ema.apply_shadow(model.unet)
8686
logger.info(f"✅ EMA weights applied (step {ema.step_count}) — better image quality")
8787

8888
elif "model_state_dict" in ckpt:

src/SD_Model.py

Lines changed: 8 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -660,12 +660,14 @@ class DDIMScheduler:
660660

661661
def __init__(
662662
self,
663-
steps: int = 1000,
664-
beta_start: float = 0.00085,
665-
beta_end: float = 0.012,
666-
schedule: str = "scaled_linear",
663+
steps: int = 1000,
664+
beta_start: float = 0.00085,
665+
beta_end: float = 0.012,
666+
schedule: str = "scaled_linear",
667+
clamp_pred_x0: bool = False,
667668
):
668669
self.num_train_timesteps = steps
670+
self.clamp_pred_x0 = clamp_pred_x0
669671

670672
if schedule == "scaled_linear":
671673
betas = torch.linspace(beta_start ** 0.5, beta_end ** 0.5, steps) ** 2
@@ -678,10 +680,6 @@ def __init__(
678680
self.timesteps: Optional[torch.Tensor] = None
679681
self.num_inference_steps: Optional[int] = None
680682

681-
def to(self, device: torch.device):
682-
self.alphas_cumprod = self.alphas_cumprod.to(device)
683-
return self
684-
685683
def set_timesteps(self, num_steps: int, device: torch.device):
686684
"""
687685
Compute the subset of evenly-spaced timesteps used for inference.
@@ -733,7 +731,8 @@ def step(
733731

734732
# Step 1: Estimate clean latent x̂_0 from noisy x_t
735733
pred_x0 = (x_t - (1.0 - alpha_t).sqrt() * noise_pred) / alpha_t.sqrt()
736-
pred_x0 = pred_x0.clamp(-1.0, 1.0) # prevent extreme values propagating
734+
if self.clamp_pred_x0:
735+
pred_x0 = pred_x0.clamp(-1.0, 1.0)
737736

738737
# Step 2: Direction from x̂_0 towards x_t (diffusion "velocity")
739738
dir_xt = (1.0 - alpha_prev).sqrt() * noise_pred

src/SD_Train.py

Lines changed: 11 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -324,6 +324,7 @@ def train_epoch(
324324
uncond_text_emb=None, use_wandb=False, use_min_snr=True,
325325
min_snr_gamma=5.0, cfg_dropout=0.0,
326326
save_steps=0, ckpt_dir="checkpoints", best_loss=float("inf"),
327+
memory_format: str = "channels_last",
327328
) -> tuple[float, int]:
328329
ddp_unet.train()
329330
model.text_encoder.eval()
@@ -338,7 +339,8 @@ def train_epoch(
338339
continue
339340
try:
340341
latents = batch["pixel_values"].to(device, dtype=torch.bfloat16, non_blocking=True)
341-
latents = latents.contiguous(memory_format=torch.channels_last)
342+
if memory_format == "channels_last":
343+
latents = latents.contiguous(memory_format=torch.channels_last)
342344
ids = batch["input_ids"].to(device, non_blocking=True)
343345
mask = batch["attention_mask"].to(device, non_blocking=True)
344346

@@ -603,7 +605,8 @@ def main(rank, world_size, args):
603605

604606
noise_scheduler = DDPMScheduler(steps=1000, beta_start=0.00085, beta_end=0.012, schedule="scaled_linear")
605607
model = StableDiffusionModel(vae, text_enc, unet, noise_scheduler).to(device)
606-
model.unet = model.unet.to(memory_format=torch.channels_last)
608+
if args.memory_format == "channels_last":
609+
model.unet = model.unet.to(memory_format=torch.channels_last)
607610
noise_scheduler.to(device)
608611

609612
# ── Dataset + unconditional embedding (precomputed once) ──────────────────
@@ -681,6 +684,7 @@ def main(rank, world_size, args):
681684
use_wandb=args.use_wandb, use_min_snr=args.min_snr, min_snr_gamma=args.min_snr_gamma,
682685
cfg_dropout=args.cfg_dropout,
683686
save_steps=args.save_steps, ckpt_dir=args.ckpt_dir, best_loss=best_loss,
687+
memory_format=args.memory_format,
684688
)
685689

686690
if avg_loss < best_loss:
@@ -720,7 +724,7 @@ def main(rank, world_size, args):
720724

721725
# Data
722726
parser.add_argument("--cache_path", type=str, default="laion_hf_dataset/train")
723-
parser.add_argument("--latent_dir", type=str, default="laion_latents/laion_latents")
727+
parser.add_argument("--latent_dir", type=str, default="laion_latents")
724728
parser.add_argument("--val_size", type=int, default=500)
725729

726730
# Model
@@ -743,6 +747,10 @@ def main(rank, world_size, args):
743747
parser.add_argument("--no-min-snr", dest="min_snr", action="store_false")
744748
parser.add_argument("--min_snr_gamma", type=float, default=5.0)
745749
parser.add_argument("--cfg_dropout", type=float, default=0.05, help="CFG dropout probability.")
750+
parser.add_argument("--memory_format", type=str, default="channels_last",
751+
choices=("channels_last", "contiguous"),
752+
help="Memory format for UNet and latents. channels_last speeds up convs on NVIDIA GPUs. "
753+
"Use 'contiguous' for AMD or Apple Silicon.")
746754

747755
# Checkpointing
748756
parser.add_argument("--save_every", type=int, default=1, help="Save checkpoint every N epochs.")

src/SD_Train_v2.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1485,7 +1485,7 @@ def main(rank: int, world_size: int, args: argparse.Namespace):
14851485
# ── Data ──────────────────────────────────────────────────────────────────
14861486
parser.add_argument("--cache_path", type=str, default="laion_hf_dataset/train",
14871487
help="Path to HuggingFace dataset (Arrow format from 05_build_hf_dataset.py).")
1488-
parser.add_argument("--latent_dir", type=str, default="laion_latents/laion_latents",
1488+
parser.add_argument("--latent_dir", type=str, default="laion_latents",
14891489
help="Directory of pre-cached .npy latent files (v1 latents are reusable).")
14901490
parser.add_argument("--val_size", type=int, default=500)
14911491
parser.add_argument("--latent_fraction", type=float, default=1.0,

src/inference.py

Lines changed: 27 additions & 22 deletions
Original file line numberDiff line numberDiff line change
@@ -160,24 +160,30 @@ def load_model(checkpoint_path: str, device: torch.device):
160160

161161
@torch.no_grad()
162162
def generate(
163-
prompts: list,
163+
prompts: list,
164164
vae,
165165
text_encoder,
166166
unet,
167167
tokenizer,
168168
scheduler,
169-
device: torch.device,
170-
num_steps: int = 50,
171-
guidance_scale: float = 7.5,
172-
seed: int = 42,
173-
height: int = 512,
174-
width: int = 512,
169+
device: torch.device,
170+
num_steps: int = 50,
171+
guidance_scale: float = 7.5,
172+
seed: int = 42,
173+
height: int = 512,
174+
width: int = 512,
175+
negative_prompts: list = None,
175176
) -> list:
176177
assert height % 8 == 0 and width % 8 == 0
177178
batch_size = len(prompts)
178179
latent_h = height // 8
179180
latent_w = width // 8
180181

182+
if negative_prompts is None:
183+
negative_prompts = [""] * batch_size
184+
elif len(negative_prompts) == 1 and batch_size > 1:
185+
negative_prompts = negative_prompts * batch_size
186+
181187
# Encode text
182188
def encode(texts):
183189
tok = tokenizer(
@@ -188,7 +194,7 @@ def encode(texts):
188194
return emb.float()
189195

190196
cond_emb = encode(prompts)
191-
uncond_emb = encode([""] * batch_size)
197+
uncond_emb = encode(negative_prompts)
192198
ctx = torch.cat([uncond_emb, cond_emb], dim=0)
193199

194200
# Initial noise — generate on CPU then move (MPS generator workaround)
@@ -279,23 +285,22 @@ def main():
279285
all_images = []
280286
for i in range(0, len(prompts), args.batch_size):
281287
batch = prompts[i:i + args.batch_size]
282-
# Apply negative prompt by replacing uncond with negative embedding
283-
# (handled inside generate via the empty string default)
284288
print(f"Generating: {batch[0][:80]}")
285289
t0 = time.time()
286290
imgs = generate(
287-
prompts = batch,
288-
vae = vae,
289-
text_encoder = text_encoder,
290-
unet = unet,
291-
tokenizer = tokenizer,
292-
scheduler = scheduler,
293-
device = device,
294-
num_steps = args.steps,
295-
guidance_scale = args.guidance,
296-
seed = args.seed + i,
297-
height = args.height,
298-
width = args.width,
291+
prompts = batch,
292+
vae = vae,
293+
text_encoder = text_encoder,
294+
unet = unet,
295+
tokenizer = tokenizer,
296+
scheduler = scheduler,
297+
device = device,
298+
num_steps = args.steps,
299+
guidance_scale = args.guidance,
300+
seed = args.seed + i,
301+
height = args.height,
302+
width = args.width,
303+
negative_prompts = [args.negative] if args.negative else None,
299304
)
300305
elapsed = time.time() - t0
301306
print(f" Done in {elapsed:.1f}s ({elapsed/args.steps:.2f}s/step)")

src/model.py

Lines changed: 8 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -682,12 +682,14 @@ class DDIMScheduler:
682682

683683
def __init__(
684684
self,
685-
steps: int = 1000,
686-
beta_start: float = 0.00085,
687-
beta_end: float = 0.012,
688-
schedule: str = "scaled_linear",
685+
steps: int = 1000,
686+
beta_start: float = 0.00085,
687+
beta_end: float = 0.012,
688+
schedule: str = "scaled_linear",
689+
clamp_pred_x0: bool = False,
689690
):
690691
self.num_train_timesteps = steps
692+
self.clamp_pred_x0 = clamp_pred_x0
691693

692694
if schedule == "scaled_linear":
693695
betas = torch.linspace(beta_start ** 0.5, beta_end ** 0.5, steps) ** 2
@@ -751,7 +753,8 @@ def step(
751753

752754
# Step 1: Estimate clean latent x̂_0 from noisy x_t
753755
pred_x0 = (x_t - (1.0 - alpha_t).sqrt() * noise_pred) / alpha_t.sqrt()
754-
pred_x0 = pred_x0.clamp(-1.0, 1.0) # prevent extreme values propagating
756+
if self.clamp_pred_x0:
757+
pred_x0 = pred_x0.clamp(-1.0, 1.0) # prevent extreme values propagating
755758

756759
# Step 2: Direction from x̂_0 towards x_t (diffusion "velocity")
757760
dir_xt = (1.0 - alpha_prev).sqrt() * noise_pred

src/train.py

Lines changed: 11 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -324,6 +324,7 @@ def train_epoch(
324324
uncond_text_emb=None, use_wandb=False, use_min_snr=True,
325325
min_snr_gamma=5.0, cfg_dropout=0.0,
326326
save_steps=0, ckpt_dir="checkpoints", best_loss=float("inf"),
327+
memory_format: str = "channels_last",
327328
) -> tuple[float, int]:
328329
ddp_unet.train()
329330
model.text_encoder.eval()
@@ -337,7 +338,8 @@ def train_epoch(
337338
continue
338339
try:
339340
latents = batch["pixel_values"].to(device, dtype=torch.bfloat16, non_blocking=True)
340-
latents = latents.contiguous(memory_format=torch.channels_last)
341+
if memory_format == "channels_last":
342+
latents = latents.contiguous(memory_format=torch.channels_last)
341343
ids = batch["input_ids"].to(device, non_blocking=True)
342344
mask = batch["attention_mask"].to(device, non_blocking=True)
343345

@@ -611,7 +613,8 @@ def main(rank, world_size, args):
611613

612614
noise_scheduler = DDPMScheduler(steps=1000, beta_start=0.00085, beta_end=0.012, schedule="scaled_linear")
613615
model = StableDiffusionModel(vae, text_enc, unet, noise_scheduler).to(device)
614-
model.unet = model.unet.to(memory_format=torch.channels_last)
616+
if args.memory_format == "channels_last":
617+
model.unet = model.unet.to(memory_format=torch.channels_last)
615618
noise_scheduler.to(device)
616619

617620
# ── Dataset + unconditional embedding (precomputed once) ──────────────────
@@ -689,6 +692,7 @@ def main(rank, world_size, args):
689692
use_wandb=args.use_wandb, use_min_snr=args.min_snr, min_snr_gamma=args.min_snr_gamma,
690693
cfg_dropout=args.cfg_dropout,
691694
save_steps=args.save_steps, ckpt_dir=args.ckpt_dir, best_loss=best_loss,
695+
memory_format=args.memory_format,
692696
)
693697

694698
if avg_loss < best_loss:
@@ -728,7 +732,7 @@ def main(rank, world_size, args):
728732

729733
# Data
730734
parser.add_argument("--cache_path", type=str, default="laion_hf_dataset/train")
731-
parser.add_argument("--latent_dir", type=str, default="laion_latents/laion_latents")
735+
parser.add_argument("--latent_dir", type=str, default="laion_latents")
732736
parser.add_argument("--val_size", type=int, default=500)
733737

734738
# Model
@@ -751,6 +755,10 @@ def main(rank, world_size, args):
751755
parser.add_argument("--no-min-snr", dest="min_snr", action="store_false")
752756
parser.add_argument("--min_snr_gamma", type=float, default=5.0)
753757
parser.add_argument("--cfg_dropout", type=float, default=0.05, help="CFG dropout probability.")
758+
parser.add_argument("--memory_format", type=str, default="channels_last",
759+
choices=("channels_last", "contiguous"),
760+
help="Memory format for UNet and latents. channels_last speeds up convs on NVIDIA GPUs. "
761+
"Use 'contiguous' for AMD or Apple Silicon.")
754762

755763
# Checkpointing
756764
parser.add_argument("--save_every", type=int, default=1)

0 commit comments

Comments
 (0)