Team Ai
Apppublic

dinhdat1110/diffusion-model

sourceHugging Faceupdated 3y agoView on Hugging Face
0likes
__main__.py171 linesDownload Raw Back to train
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