diff --git a/sample.py b/sample.py index a4152afd..acf9b67b 100644 --- a/sample.py +++ b/sample.py @@ -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) @@ -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 @@ -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) diff --git a/train.py b/train.py index 7cfee808..8eeae3e5 100644 --- a/train.py +++ b/train.py @@ -26,6 +26,7 @@ import argparse import logging import os +import wandb from models import DiT_models from diffusion import create_diffusion @@ -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) @@ -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()):,}") @@ -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() @@ -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 @@ -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 ... @@ -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)