Skip to content
This repository was archived by the owner on Aug 6, 2025. It is now read-only.
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 4 additions & 3 deletions sample.py
Original file line number Diff line number Diff line change
Expand Up @@ -40,11 +40,11 @@ def main(args):
state_dict = find_model(ckpt_path)
model.load_state_dict(state_dict)
model.eval() # important!
diffusion = create_diffusion(str(args.num_sampling_steps))
diffusion = create_diffusion(str(args.num_sampling_steps), diffusion_steps=args.diffusion_steps)
vae = AutoencoderKL.from_pretrained(f"stabilityai/sd-vae-ft-{args.vae}").to(device)

# Labels to condition the model with (feel free to change):
class_labels = [207, 360, 387, 974, 88, 979, 417, 279]
class_labels = [207, 0, 0, 0, 0, 0, 0, 0]

# Create sampling noise:
n = len(class_labels)
Expand All @@ -59,7 +59,7 @@ def main(args):

# Sample images:
samples = diffusion.p_sample_loop(
model.forward_with_cfg, z.shape, z, clip_denoised=False, model_kwargs=model_kwargs, progress=True, device=device
model.forward_with_cfg, z.shape, z, clip_denoised=True, model_kwargs=model_kwargs, progress=True, device=device
)
samples, _ = samples.chunk(2, dim=0) # Remove null class samples
samples = vae.decode(samples / 0.18215).sample
Expand All @@ -79,5 +79,6 @@ def main(args):
parser.add_argument("--seed", type=int, default=0)
parser.add_argument("--ckpt", type=str, default=None,
help="Optional path to a DiT checkpoint (default: auto-download a pre-trained DiT-XL/2 model).")
parser.add_argument("--diffusion-steps", type=int, default=50)
args = parser.parse_args()
main(args)
51 changes: 49 additions & 2 deletions train.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,7 @@
import argparse
import logging
import os
import wandb

from models import DiT_models
from diffusion import create_diffusion
Expand Down Expand Up @@ -133,6 +134,11 @@ def main(args):
os.makedirs(checkpoint_dir, exist_ok=True)
logger = create_logger(experiment_dir)
logger.info(f"Experiment directory created at {experiment_dir}")
if args.wandb_logging:
wandb.init(
project=args.wandb_project,
name=args.wandb_name,
)
else:
logger = create_logger(None)

Expand All @@ -147,7 +153,7 @@ def main(args):
ema = deepcopy(model).to(device) # Create an EMA of the model for use after training
requires_grad(ema, False)
model = DDP(model.to(device), device_ids=[rank])
diffusion = create_diffusion(timestep_respacing="") # default: 1000 steps, linear noise schedule
diffusion = create_diffusion(timestep_respacing="", diffusion_steps=args.diffusion_steps) # default: 1000 steps, linear noise schedule
vae = AutoencoderKL.from_pretrained(f"stabilityai/sd-vae-ft-{args.vae}").to(device)
logger.info(f"DiT Parameters: {sum(p.numel() for p in model.parameters()):,}")

Expand Down Expand Up @@ -208,7 +214,7 @@ def main(args):
opt.zero_grad()
loss.backward()
opt.step()
update_ema(ema, model.module)
update_ema(ema, model.module, decay=args.ema_decay)

# Log loss values:
running_loss += loss.item()
Expand All @@ -224,6 +230,7 @@ def main(args):
dist.all_reduce(avg_loss, op=dist.ReduceOp.SUM)
avg_loss = avg_loss.item() / dist.get_world_size()
logger.info(f"(step={train_steps:07d}) Train Loss: {avg_loss:.4f}, Train Steps/Sec: {steps_per_sec:.2f}")
wandb.log({"Train Loss": avg_loss, "Train Steps/Sec": steps_per_sec}, step=train_steps)
# Reset monitoring variables:
running_loss = 0
log_steps = 0
Expand All @@ -243,6 +250,39 @@ def main(args):
logger.info(f"Saved checkpoint to {checkpoint_path}")
dist.barrier()

# Sample generated imaged from current model:
if train_steps % args.sample_every == 0:
model.eval()
with torch.no_grad():
n_samples = 4
z = torch.randn(
n_samples * 2, 4,
latent_size, latent_size,
device=device
)
# TODO: replace 1 by nb_classes
y = torch.randint(0, 1, (n_samples,), device=device)
y_null = torch.tensor([1000] * n_samples, device=device)
y_sample = torch.cat([y, y_null], 0)
model_kwargs = dict(y=y_sample, cfg_scale=args.sample_cfg_scale)
samples = diffusion.p_sample_loop(
model.module.forward_with_cfg,
z.shape, z,
clip_denoised=True,
model_kwargs=model_kwargs,
progress=False,
device=device
)
samples, _ = samples.chunk(2, dim=0)
samples = vae.decode(samples / 0.18215).sample
samples = (samples.clamp(-1, 1) + 1) / 2 # Normalize to [0, 1]
if args.wandb_logging:
wandb.log(
{"Generated Images": [wandb.Image(img) for img in samples]},
step=train_steps
)
model.train()

model.eval() # important! This disables randomized embedding dropout
# do any sampling/FID calculation/etc. with ema (or model) in eval mode ...

Expand All @@ -265,5 +305,12 @@ def main(args):
parser.add_argument("--num-workers", type=int, default=4)
parser.add_argument("--log-every", type=int, default=100)
parser.add_argument("--ckpt-every", type=int, default=50_000)
parser.add_argument("--wandb-logging", type=bool, default=False)
parser.add_argument("--wandb-project", type=str, default="DiT")
parser.add_argument("--wandb-name", type=str, default="fashion-dataset")
parser.add_argument("--sample-every", type=int, default=100)
parser.add_argument("--sample-cfg-scale", type=float, default=4.0, help="Classifier-free guidance scale.")
parser.add_argument("--diffusion-steps", type=int, default=1000) # 100 for small experiments
parser.add_argument("--ema-decay", type=float, default=0.9999) # 0.999 for small experiments
args = parser.parse_args()
main(args)