Anonymous-123/ImageNet-Editing
1
1"""2Train a noised image classifier on ImageNet.3"""4 5import argparse6import os7 8import blobfile as bf9import torch as th10import torch.distributed as dist11import torch.nn.functional as F12from torch.nn.parallel.distributed import DistributedDataParallel as DDP13from torch.optim import AdamW14 15from guided_diffusion import dist_util, logger16from guided_diffusion.fp16_util import MixedPrecisionTrainer17from guided_diffusion.image_datasets import load_data18from guided_diffusion.resample import create_named_schedule_sampler19from guided_diffusion.script_util import (20 add_dict_to_argparser,21 args_to_dict,22 classifier_and_diffusion_defaults,23 create_classifier_and_diffusion,24)25from guided_diffusion.train_util import parse_resume_step_from_filename, log_loss_dict26 27 28def main():29 args = create_argparser().parse_args()30 31 dist_util.setup_dist()32 logger.configure()33 34 logger.log("creating model and diffusion...")35 model, diffusion = create_classifier_and_diffusion(36 **args_to_dict(args, classifier_and_diffusion_defaults().keys())37 )38 model.to(dist_util.dev())39 if args.noised:40 schedule_sampler = create_named_schedule_sampler(41 args.schedule_sampler, diffusion42 )43 44 resume_step = 045 if args.resume_checkpoint:46 resume_step = parse_resume_step_from_filename(args.resume_checkpoint)47 if dist.get_rank() == 0:48 logger.log(49 f"loading model from checkpoint: {args.resume_checkpoint}... at {resume_step} step"50 )51 model.load_state_dict(52 dist_util.load_state_dict(53 args.resume_checkpoint, map_location=dist_util.dev()54 )55 )56 57 # Needed for creating correct EMAs and fp16 parameters.58 dist_util.sync_params(model.parameters())59 60 mp_trainer = MixedPrecisionTrainer(61 model=model, use_fp16=args.classifier_use_fp16, initial_lg_loss_scale=16.062 )63 64 model = DDP(65 model,66 device_ids=[dist_util.dev()],67 output_device=dist_util.dev(),68 broadcast_buffers=False,69 bucket_cap_mb=128,70 find_unused_parameters=False,71 )72 73 logger.log("creating data loader...")74 data = load_data(75 data_dir=args.data_dir,76 batch_size=args.batch_size,77 image_size=args.image_size,78 class_cond=True,79 random_crop=True,80 )81 if args.val_data_dir:82 val_data = load_data(83 data_dir=args.val_data_dir,84 batch_size=args.batch_size,85 image_size=args.image_size,86 class_cond=True,87 )88 else:89 val_data = None90 91 logger.log(f"creating optimizer...")92 opt = AdamW(mp_trainer.master_params, lr=args.lr, weight_decay=args.weight_decay)93 if args.resume_checkpoint:94 opt_checkpoint = bf.join(95 bf.dirname(args.resume_checkpoint), f"opt{resume_step:06}.pt"96 )97 logger.log(f"loading optimizer state from checkpoint: {opt_checkpoint}")98 opt.load_state_dict(99 dist_util.load_state_dict(opt_checkpoint, map_location=dist_util.dev())100 )101 102 logger.log("training classifier model...")103 104 def forward_backward_log(data_loader, prefix="train"):105 batch, extra = next(data_loader)106 labels = extra["y"].to(dist_util.dev())107 108 batch = batch.to(dist_util.dev())109 # Noisy images110 if args.noised:111 t, _ = schedule_sampler.sample(batch.shape[0], dist_util.dev())112 batch = diffusion.q_sample(batch, t)113 else:114 t = th.zeros(batch.shape[0], dtype=th.long, device=dist_util.dev())115 116 for i, (sub_batch, sub_labels, sub_t) in enumerate(117 split_microbatches(args.microbatch, batch, labels, t)118 ):119 logits = model(sub_batch, timesteps=sub_t)120 loss = F.cross_entropy(logits, sub_labels, reduction="none")121 122 losses = {}123 losses[f"{prefix}_loss"] = loss.detach()124 losses[f"{prefix}_acc@1"] = compute_top_k(125 logits, sub_labels, k=1, reduction="none"126 )127 losses[f"{prefix}_acc@5"] = compute_top_k(128 logits, sub_labels, k=5, reduction="none"129 )130 log_loss_dict(diffusion, sub_t, losses)131 del losses132 loss = loss.mean()133 if loss.requires_grad:134 if i == 0:135 mp_trainer.zero_grad()136 mp_trainer.backward(loss * len(sub_batch) / len(batch))137 138 for step in range(args.iterations - resume_step):139 logger.logkv("step", step + resume_step)140 logger.logkv(141 "samples",142 (step + resume_step + 1) * args.batch_size * dist.get_world_size(),143 )144 if args.anneal_lr:145 set_annealed_lr(opt, args.lr, (step + resume_step) / args.iterations)146 forward_backward_log(data)147 mp_trainer.optimize(opt)148 if val_data is not None and not step % args.eval_interval:149 with th.no_grad():150 with model.no_sync():151 model.eval()152 forward_backward_log(val_data, prefix="val")153 model.train()154 if not step % args.log_interval:155 logger.dumpkvs()156 if (157 step158 and dist.get_rank() == 0159 and not (step + resume_step) % args.save_interval160 ):161 logger.log("saving model...")162 save_model(mp_trainer, opt, step + resume_step)163 164 if dist.get_rank() == 0:165 logger.log("saving model...")166 save_model(mp_trainer, opt, step + resume_step)167 dist.barrier()168 169 170def set_annealed_lr(opt, base_lr, frac_done):171 lr = base_lr * (1 - frac_done)172 for param_group in opt.param_groups:173 param_group["lr"] = lr174 175 176def save_model(mp_trainer, opt, step):177 if dist.get_rank() == 0:178 th.save(179 mp_trainer.master_params_to_state_dict(mp_trainer.master_params),180 os.path.join(logger.get_dir(), f"model{step:06d}.pt"),181 )182 th.save(opt.state_dict(), os.path.join(logger.get_dir(), f"opt{step:06d}.pt"))183 184 185def compute_top_k(logits, labels, k, reduction="mean"):186 _, top_ks = th.topk(logits, k, dim=-1)187 if reduction == "mean":188 return (top_ks == labels[:, None]).float().sum(dim=-1).mean().item()189 elif reduction == "none":190 return (top_ks == labels[:, None]).float().sum(dim=-1)191 192 193def split_microbatches(microbatch, *args):194 bs = len(args[0])195 if microbatch == -1 or microbatch >= bs:196 yield tuple(args)197 else:198 for i in range(0, bs, microbatch):199 yield tuple(x[i : i + microbatch] if x is not None else None for x in args)200 201 202def create_argparser():203 defaults = dict(204 data_dir="",205 val_data_dir="",206 noised=True,207 iterations=150000,208 lr=3e-4,209 weight_decay=0.0,210 anneal_lr=False,211 batch_size=4,212 microbatch=-1,213 schedule_sampler="uniform",214 resume_checkpoint="",215 log_interval=10,216 eval_interval=5,217 save_interval=10000,218 )219 defaults.update(classifier_and_diffusion_defaults())220 parser = argparse.ArgumentParser()221 add_dict_to_argparser(parser, defaults)222 return parser223 224 225if __name__ == "__main__":226 main()227 