CoolFace
Apppublic

Anonymous-123/ImageNet-Editing

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