| import torchvision |
| import os |
| import errno |
| import shutil |
| import argparse |
| from networks import TwoBranchModel,Unet |
| from diffusion_pytorch import GaussianDiffusion, Trainer |
| import torch, warnings |
|
|
| from pytorch_lightning.callbacks import Callback |
| warnings.filterwarnings("ignore") |
|
|
|
|
| class DebugDataloaderCallback(Callback): |
| |
| def __init__(self): |
| super().__init__() |
| self.counter = 0 |
|
|
| def on_train_start(self, trainer, pl_module): |
| self.counter += 1 |
| if (self.counter + 1 ) % 10 == 0: |
| trainer.train_dataloader.dataset.update_chunk() |
|
|
|
|
|
|
| def create_folder(path): |
| try: |
| os.mkdir(path) |
| except OSError as exc: |
| if exc.errno != errno.EEXIST: |
| raise |
| pass |
|
|
|
|
| def del_folder(path): |
| try: |
| shutil.rmtree(path) |
| except OSError as exc: |
| pass |
|
|
|
|
| create = 0 |
|
|
| if create: |
| trainset = torchvision.datasets.CIFAR10( |
| root='./data', train=True, download=True) |
| root = './root_cifar10/' |
| del_folder(root) |
| create_folder(root) |
|
|
| for i in range(10): |
| lable_root = root + str(i) + '/' |
| create_folder(lable_root) |
|
|
| for idx in range(len(trainset)): |
| img, label = trainset[idx] |
| print(idx) |
| img.save(root + str(label) + '/' + str(idx) + '.png') |
|
|
|
|
| parser = argparse.ArgumentParser() |
| parser.add_argument('--time_steps', default=50, type=int) |
| parser.add_argument('--train_steps', default=700000, type=int) |
| parser.add_argument('--save_folder', default=None, type=str) |
|
|
| parser.add_argument('--load_path', default=None, type=str) |
| parser.add_argument('--data_path', default='./root_cifar10/', type=str) |
| parser.add_argument('--fade_routine', default='Random_Incremental', type=str) |
| parser.add_argument('--sampling_routine', default='x0_step_down', type=str) |
| parser.add_argument('--discrete', action="store_true") |
| parser.add_argument('--remove_time_embed', action="store_true") |
| parser.add_argument('--residual', action="store_true") |
| parser.add_argument('--tag', default='', type=str) |
| parser.add_argument('--accelerate_factor', default=4, help="4 | 8", type=int) |
|
|
|
|
| parser.add_argument('--normalizer', default='mean_std', type=str) |
|
|
| parser.add_argument('--mode', default='train', type=str) |
| parser.add_argument('--example_frequency_img', default=None, type=str) |
| |
| |
| parser.add_argument('--kernel_std', default=0.1, type=float) |
|
|
| parser.add_argument('--dataset', default='brain', type=str) |
| parser.add_argument('--domain', default=None, type=str) |
| parser.add_argument('--aux_modality', default=None, type=str) |
| parser.add_argument('--deviceid', default=0, type=int) |
| parser.add_argument('--num_channels', default=1, type=int) |
| parser.add_argument('--train_bs', default=24, type=int) |
| parser.add_argument('--diffusion_type', default='twobranch_fade', type=str) |
| parser.add_argument('--debug', action="store_true") |
| parser.add_argument('--image_size', default=128) |
| parser.add_argument('--loss_type', default='l1', type=str) |
|
|
| args = parser.parse_args() |
| print(args) |
| os.environ["CUDA_VISIBLE_DEVICES"] = str(args.deviceid) |
|
|
| image_channels = 1 |
|
|
| diffusion_type = args.diffusion_type |
| |
| model_name = diffusion_type.split("_")[0] |
|
|
| save_and_sample_every = 1000 |
|
|
| if args.debug: |
| args.train_steps = 100 |
| args.time_steps = 5 |
|
|
| model = None |
|
|
|
|
| if isinstance(args.image_size, str): |
| length = len(args.image_size.split(",")) |
| if length == 1: |
| args.image_size = (int(args.image_size), int(args.image_size)) |
| elif length == 2: |
| args.image_size = (int(args.image_size.split(",")[0]), int(args.image_size.split(",")[1])) |
| else: |
| args.image_size = (args.image_size, args.image_size) |
|
|
|
|
|
|
| if model_name == "unet": |
| model = Unet(resolution=args.image_size[0], |
| in_channels=1, |
| out_ch=1, |
| ch=128, |
| ch_mult=(1, 2, 2, 2), |
| num_res_blocks=2, |
| attn_resolutions=(16,), |
| dropout=0.1).cuda() |
|
|
| elif model_name == "twounet": |
| model = TwoBranchNewModel(resolution=args.image_size[0], |
| in_channels=1, |
| out_ch=1, |
| ch=128, |
| ch_mult=(1, 2, 2, 2), |
| num_res_blocks=3, |
| attn_resolutions=(16,), |
| dropout=0.1).cuda() |
|
|
|
|
| elif model_name == "twobranch": |
|
|
| base_num_every_group = 2 |
| num_features = 64 |
| act = "PReLU" |
| num_channels = 1 |
|
|
| from networks.networks_fsm.mynet import TwoBranch as TwoBranchModel |
|
|
|
|
| model = TwoBranchModel( |
| num_features, act, base_num_every_group, num_channels |
| ).cuda() |
|
|
| fp16 = False |
|
|
|
|
| n_parameters = sum(p.numel() for p in model.parameters() if p.requires_grad) |
| print('number of params: %.2f M' % (n_parameters / 1024 / 1024)) |
|
|
|
|
| diffusion = GaussianDiffusion( |
| diffusion_type, |
| model, |
| image_size=args.image_size[0], |
| channels=image_channels, |
| device_of_kernel='cuda', |
| timesteps=args.time_steps, |
| loss_type=args.loss_type, |
| kernel_std=args.kernel_std, |
| fade_routine=args.fade_routine, |
| sampling_routine=args.sampling_routine, |
| discrete=args.discrete, |
| accelerate_factor=args.accelerate_factor, |
| fp16=fp16, |
| normalizer=args.normalizer, |
| example_frequency_img=args.example_frequency_img, |
| ).cuda() |
|
|
|
|
| diffusion = torch.nn.DataParallel(diffusion, device_ids=range(torch.cuda.device_count())) |
|
|
| print("=== train_steps:", args.train_steps) |
| os.makedirs(args.save_folder, exist_ok=True) |
|
|
| if args.debug: |
| args.save_folder = args.save_folder + "_debug" |
| else: |
| args.save_folder = args.save_folder + f"_{args.tag}" |
| save_and_sample_every = 500 |
|
|
|
|
| |
| name = args.save_folder.split("/")[-1] |
| number = os.listdir(args.save_folder.rstrip(name)).__len__() |
| if args.mode == "test": |
| number = "test_" + str(number) |
|
|
| args.save_folder = os.path.join(args.save_folder.rstrip(name), f"{number}_" + name) |
|
|
| |
| os.makedirs(args.save_folder, exist_ok=True) |
|
|
|
|
|
|
| print("SAVE FOLDER: ", args.save_folder) |
|
|
| trainer = Trainer( |
| diffusion, |
| args.data_path, |
| mode = args.mode, |
| norm = args.normalizer, |
| image_size=args.image_size, |
| train_batch_size=args.train_bs, |
| train_lr= 1e-4, |
| train_num_steps=args.train_steps, |
| gradient_accumulate_every=1, |
| ema_decay=0.995, |
| save_and_sample_every=save_and_sample_every, |
| fp16=fp16, |
| results_folder=args.save_folder, |
| load_path=args.load_path, |
| dataset=args.dataset, |
| domain=args.domain, |
| aux_modality=args.aux_modality, |
| debug=args.debug, |
| num_channels=args.num_channels |
| |
| |
| ) |
|
|
|
|
| if args.mode == "train": |
| trainer.train() |
|
|
| elif args.mode == "test": |
| |
| trainer.test_loader('x0_step_down_fre') |
|
|
|
|
|
|
|
|
|
|