Anonymous-123/ImageNet-Editing
1
1import copy2import functools3import os4 5import blobfile as bf6import torch as th7import torch.distributed as dist8from torch.nn.parallel.distributed import DistributedDataParallel as DDP9from torch.optim import AdamW10 11from . import dist_util, logger12from .fp16_util import MixedPrecisionTrainer13from .nn import update_ema14from .resample import LossAwareSampler, UniformSampler15 16# For ImageNet experiments, this was a good default value.17# We found that the lg_loss_scale quickly climbed to18# 20-21 within the first ~1K steps of training.19INITIAL_LOG_LOSS_SCALE = 20.020 21 22class TrainLoop:23 def __init__(24 self,25 *,26 model,27 diffusion,28 data,29 batch_size,30 microbatch,31 lr,32 ema_rate,33 log_interval,34 save_interval,35 resume_checkpoint,36 use_fp16=False,37 fp16_scale_growth=1e-3,38 schedule_sampler=None,39 weight_decay=0.0,40 lr_anneal_steps=0,41 ):42 self.model = model43 self.diffusion = diffusion44 self.data = data45 self.batch_size = batch_size46 self.microbatch = microbatch if microbatch > 0 else batch_size47 self.lr = lr48 self.ema_rate = (49 [ema_rate]50 if isinstance(ema_rate, float)51 else [float(x) for x in ema_rate.split(",")]52 )53 self.log_interval = log_interval54 self.save_interval = save_interval55 self.resume_checkpoint = resume_checkpoint56 self.use_fp16 = use_fp1657 self.fp16_scale_growth = fp16_scale_growth58 self.schedule_sampler = schedule_sampler or UniformSampler(diffusion)59 self.weight_decay = weight_decay60 self.lr_anneal_steps = lr_anneal_steps61 62 self.step = 063 self.resume_step = 064 self.global_batch = self.batch_size * dist.get_world_size()65 66 self.sync_cuda = th.cuda.is_available()67 68 self._load_and_sync_parameters()69 self.mp_trainer = MixedPrecisionTrainer(70 model=self.model,71 use_fp16=self.use_fp16,72 fp16_scale_growth=fp16_scale_growth,73 )74 75 self.opt = AdamW(76 self.mp_trainer.master_params, lr=self.lr, weight_decay=self.weight_decay77 )78 if self.resume_step:79 self._load_optimizer_state()80 # Model was resumed, either due to a restart or a checkpoint81 # being specified at the command line.82 self.ema_params = [83 self._load_ema_parameters(rate) for rate in self.ema_rate84 ]85 else:86 self.ema_params = [87 copy.deepcopy(self.mp_trainer.master_params)88 for _ in range(len(self.ema_rate))89 ]90 91 if th.cuda.is_available():92 self.use_ddp = True93 self.ddp_model = DDP(94 self.model,95 device_ids=[dist_util.dev()],96 output_device=dist_util.dev(),97 broadcast_buffers=False,98 bucket_cap_mb=128,99 find_unused_parameters=False,100 )101 else:102 if dist.get_world_size() > 1:103 logger.warn(104 "Distributed training requires CUDA. "105 "Gradients will not be synchronized properly!"106 )107 self.use_ddp = False108 self.ddp_model = self.model109 110 def _load_and_sync_parameters(self):111 resume_checkpoint = find_resume_checkpoint() or self.resume_checkpoint112 113 if resume_checkpoint:114 self.resume_step = parse_resume_step_from_filename(resume_checkpoint)115 if dist.get_rank() == 0:116 logger.log(f"loading model from checkpoint: {resume_checkpoint}...")117 self.model.load_state_dict(118 dist_util.load_state_dict(119 resume_checkpoint, map_location=dist_util.dev()120 )121 )122 123 dist_util.sync_params(self.model.parameters())124 125 def _load_ema_parameters(self, rate):126 ema_params = copy.deepcopy(self.mp_trainer.master_params)127 128 main_checkpoint = find_resume_checkpoint() or self.resume_checkpoint129 ema_checkpoint = find_ema_checkpoint(main_checkpoint, self.resume_step, rate)130 if ema_checkpoint:131 if dist.get_rank() == 0:132 logger.log(f"loading EMA from checkpoint: {ema_checkpoint}...")133 state_dict = dist_util.load_state_dict(134 ema_checkpoint, map_location=dist_util.dev()135 )136 ema_params = self.mp_trainer.state_dict_to_master_params(state_dict)137 138 dist_util.sync_params(ema_params)139 return ema_params140 141 def _load_optimizer_state(self):142 main_checkpoint = find_resume_checkpoint() or self.resume_checkpoint143 opt_checkpoint = bf.join(144 bf.dirname(main_checkpoint), f"opt{self.resume_step:06}.pt"145 )146 if bf.exists(opt_checkpoint):147 logger.log(f"loading optimizer state from checkpoint: {opt_checkpoint}")148 state_dict = dist_util.load_state_dict(149 opt_checkpoint, map_location=dist_util.dev()150 )151 self.opt.load_state_dict(state_dict)152 153 def run_loop(self):154 while (155 not self.lr_anneal_steps156 or self.step + self.resume_step < self.lr_anneal_steps157 ):158 batch, cond = next(self.data)159 self.run_step(batch, cond)160 if self.step % self.log_interval == 0:161 logger.dumpkvs()162 if self.step % self.save_interval == 0:163 self.save()164 # Run for a finite amount of time in integration tests.165 if os.environ.get("DIFFUSION_TRAINING_TEST", "") and self.step > 0:166 return167 self.step += 1168 # Save the last checkpoint if it wasn't already saved.169 if (self.step - 1) % self.save_interval != 0:170 self.save()171 172 def run_step(self, batch, cond):173 self.forward_backward(batch, cond)174 took_step = self.mp_trainer.optimize(self.opt)175 if took_step:176 self._update_ema()177 self._anneal_lr()178 self.log_step()179 180 def forward_backward(self, batch, cond):181 self.mp_trainer.zero_grad()182 for i in range(0, batch.shape[0], self.microbatch):183 micro = batch[i : i + self.microbatch].to(dist_util.dev())184 micro_cond = {185 k: v[i : i + self.microbatch].to(dist_util.dev())186 for k, v in cond.items()187 }188 last_batch = (i + self.microbatch) >= batch.shape[0]189 t, weights = self.schedule_sampler.sample(micro.shape[0], dist_util.dev())190 191 compute_losses = functools.partial(192 self.diffusion.training_losses,193 self.ddp_model,194 micro,195 t,196 model_kwargs=micro_cond,197 )198 199 if last_batch or not self.use_ddp:200 losses = compute_losses()201 else:202 with self.ddp_model.no_sync():203 losses = compute_losses()204 205 if isinstance(self.schedule_sampler, LossAwareSampler):206 self.schedule_sampler.update_with_local_losses(207 t, losses["loss"].detach()208 )209 210 loss = (losses["loss"] * weights).mean()211 log_loss_dict(212 self.diffusion, t, {k: v * weights for k, v in losses.items()}213 )214 self.mp_trainer.backward(loss)215 216 def _update_ema(self):217 for rate, params in zip(self.ema_rate, self.ema_params):218 update_ema(params, self.mp_trainer.master_params, rate=rate)219 220 def _anneal_lr(self):221 if not self.lr_anneal_steps:222 return223 frac_done = (self.step + self.resume_step) / self.lr_anneal_steps224 lr = self.lr * (1 - frac_done)225 for param_group in self.opt.param_groups:226 param_group["lr"] = lr227 228 def log_step(self):229 logger.logkv("step", self.step + self.resume_step)230 logger.logkv("samples", (self.step + self.resume_step + 1) * self.global_batch)231 232 def save(self):233 def save_checkpoint(rate, params):234 state_dict = self.mp_trainer.master_params_to_state_dict(params)235 if dist.get_rank() == 0:236 logger.log(f"saving model {rate}...")237 if not rate:238 filename = f"model{(self.step+self.resume_step):06d}.pt"239 else:240 filename = f"ema_{rate}_{(self.step+self.resume_step):06d}.pt"241 with bf.BlobFile(bf.join(get_blob_logdir(), filename), "wb") as f:242 th.save(state_dict, f)243 244 save_checkpoint(0, self.mp_trainer.master_params)245 for rate, params in zip(self.ema_rate, self.ema_params):246 save_checkpoint(rate, params)247 248 if dist.get_rank() == 0:249 with bf.BlobFile(250 bf.join(get_blob_logdir(), f"opt{(self.step+self.resume_step):06d}.pt"),251 "wb",252 ) as f:253 th.save(self.opt.state_dict(), f)254 255 dist.barrier()256 257 258def parse_resume_step_from_filename(filename):259 """260 Parse filenames of the form path/to/modelNNNNNN.pt, where NNNNNN is the261 checkpoint's number of steps.262 """263 split = filename.split("model")264 if len(split) < 2:265 return 0266 split1 = split[-1].split(".")[0]267 try:268 return int(split1)269 except ValueError:270 return 0271 272 273def get_blob_logdir():274 # You can change this to be a separate path to save checkpoints to275 # a blobstore or some external drive.276 return logger.get_dir()277 278 279def find_resume_checkpoint():280 # On your infrastructure, you may want to override this to automatically281 # discover the latest checkpoint on your blob storage, etc.282 return None283 284 285def find_ema_checkpoint(main_checkpoint, step, rate):286 if main_checkpoint is None:287 return None288 filename = f"ema_{rate}_{(step):06d}.pt"289 path = bf.join(bf.dirname(main_checkpoint), filename)290 if bf.exists(path):291 return path292 return None293 294 295def log_loss_dict(diffusion, ts, losses):296 for key, values in losses.items():297 logger.logkv_mean(key, values.mean().item())298 # Log the quantiles (four quartiles, in particular).299 for sub_t, sub_loss in zip(ts.cpu().numpy(), values.detach().cpu().numpy()):300 quartile = int(4 * sub_t / diffusion.num_timesteps)301 logger.logkv_mean(f"{key}_q{quartile}", sub_loss)302 