qic999's picture
Upload folder using huggingface_hub
28e6f98 verified
Raw
History Blame Contribute Delete
7.09 kB
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')