CoolFace
Apppublic

Anonymous-123/ImageNet-Editing

sourceHugging Facecreativeml-openrail-mupdated 4y agoView on Hugging Face
1likes
train_util.py302 linesDownload Raw Back to guided_diffusion
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