dinhdat1110/diffusion-model
0
1from pytorch_lightning.loggers import WandbLogger2import diffusion3import torch4import wandb5import pytorch_lightning as pl6import argparse7import os8 9torch.multiprocessing.set_sharing_strategy('file_system')10 11 12def main():13 # PARSERs14 parser = argparse.ArgumentParser()15 parser.add_argument(16 '--dataset', '-d', type=str, default='mnist',17 help='choose dataset'18 )19 parser.add_argument(20 '--data_dir', '-dd', type=str, default='./data/',21 help='model name'22 )23 parser.add_argument(24 '--mode', type=str, default='ddim',25 help='sampling mode'26 )27 parser.add_argument(28 '--max_epochs', '-me', type=int, default=200,29 help='max epoch'30 )31 parser.add_argument(32 '--batch_size', '-bs', type=int, default=32,33 help='batch size'34 )35 parser.add_argument(36 '--train_ratio', '-tr', type=float, default=0.99,37 help='batch size'38 )39 parser.add_argument(40 '--timesteps', '-ts', type=int, default=1000,41 help='max timesteps diffusion'42 )43 parser.add_argument(44 '--max_batch_size', '-mbs', type=int, default=32,45 help='max batch size'46 )47 parser.add_argument(48 '--lr', '-l', type=float, default=1e-4,49 help='learning rate'50 )51 parser.add_argument(52 '--num_workers', '-nw', type=int, default=4,53 help='number of workers'54 )55 parser.add_argument(56 '--seed', '-s', type=int, default=42,57 help='seed'58 )59 parser.add_argument(60 '--name', '-n', type=str, default=None,61 help='name of the experiment'62 )63 parser.add_argument(64 '--pbar', action='store_true',65 help='progress bar'66 )67 parser.add_argument(68 '--precision', '-p', type=str, default='32',69 help='numerical precision'70 )71 parser.add_argument(72 '--sample_per_epochs', '-spe', type=int, default=25,73 help='sample every n epochs'74 )75 parser.add_argument(76 '--n_samples', '-ns', type=int, default=4,77 help='number of workers'78 )79 parser.add_argument(80 '--monitor', '-m', type=str, default='val_loss',81 help='callbacks monitor'82 )83 parser.add_argument(84 '--wandb', '-wk', type=str, default=None,85 help='wandb API key'86 )87 88 args = parser.parse_args()89 90 # SEED91 pl.seed_everything(args.seed, workers=True)92 93 # WANDB (OPTIONAL)94 if args.wandb is not None:95 wandb.login(key=args.wandb) # API KEY96 name = args.name or f"diffusion-{args.max_epochs}-{args.batch_size}-{args.lr}"97 logger = WandbLogger(98 project="diffusion-model",99 name=name,100 log_model=False101 )102 else:103 logger = None104 105 # DATAMODULE106 if args.dataset == "mnist":107 DATAMODULE = diffusion.MNISTDataModule108 img_dim = 32109 num_classes = 10110 elif args.dataset == "cifar10":111 DATAMODULE = diffusion.CIFAR10DataModule112 img_dim = 32113 num_classes = 10114 elif args.dataset == "celeba":115 DATAMODULE = diffusion.CelebADataModule116 img_dim = 64117 num_classes = None118 119 datamodule = DATAMODULE(120 data_dir=args.data_dir,121 batch_size=args.batch_size,122 num_workers=args.num_workers,123 seed=args.seed,124 train_ratio=args.train_ratio,125 img_dim=img_dim126 )127 128 # MODEL129 in_channels = 1 if args.dataset == "mnist" else 3130 model = diffusion.DiffusionModel(131 lr=args.lr,132 in_channels=in_channels,133 sample_per_epochs=args.sample_per_epochs,134 max_timesteps=args.timesteps,135 dim=img_dim,136 num_classes=num_classes,137 n_samples=args.n_samples,138 mode=args.mode139 )140 141 # CALLBACK142 root_path = os.path.join(os.getcwd(), "checkpoints")143 callback = diffusion.ModelCallback(144 root_path=root_path,145 ckpt_monitor=args.monitor146 )147 148 # STRATEGY149 strategy = 'ddp_find_unused_parameters_true' if torch.cuda.is_available() else 'auto'150 151 # TRAINER152 trainer = pl.Trainer(153 default_root_dir=root_path,154 logger=logger,155 callbacks=callback.get_callback(),156 gradient_clip_val=0.5,157 max_epochs=args.max_epochs,158 enable_progress_bar=args.pbar,159 deterministic=False,160 precision=args.precision,161 strategy=strategy,162 accumulate_grad_batches=max(int(args.max_batch_size / args.batch_size), 1)163 )164 165 # FIT MODEL166 trainer.fit(model=model, datamodule=datamodule)167 168 169if __name__ == '__main__':170 main()171 