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) # specific arguments # parser.add_argument('--initial_mask', default=11, type=int) 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 # diffusion_type = "twobranch_fade" # model_degradation # fade | kspace model_name = diffusion_type.split("_")[0] # unet | twobranch 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() # Drop out used to be 0.1 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], # Used to be 32 channels=image_channels, device_of_kernel='cuda', timesteps=args.time_steps, loss_type=args.loss_type, #$'l1', 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 # if os.path.exists(args.save_folder): 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) # create the folder and parent folders 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, # Used to be 32 train_batch_size=args.train_bs, train_lr= 1e-4, # 2e-5 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 # accelerator="gpu", # callbacks=[DebugDataloaderCallback()], ) if args.mode == "train": trainer.train() elif args.mode == "test": # ['default', 'x0_step_down', 'x0_step_down_fre', "fre_progressive"]: trainer.test_loader('x0_step_down_fre')