CoolFace
Apppublic

Anonymous-123/ImageNet-Editing

sourceHugging Facecreativeml-openrail-mupdated 4y agoView on Hugging Face
1likes
image_nll.py97 linesDownload Raw Back to scripts
1"""2Approximate the bits/dimension for an image model.3"""4 5import argparse6import os7 8import numpy as np9import torch.distributed as dist10 11from guided_diffusion import dist_util, logger12from guided_diffusion.image_datasets import load_data13from guided_diffusion.script_util import (14    model_and_diffusion_defaults,15    create_model_and_diffusion,16    add_dict_to_argparser,17    args_to_dict,18)19 20 21def main():22    args = create_argparser().parse_args()23 24    dist_util.setup_dist()25    logger.configure()26 27    logger.log("creating model and diffusion...")28    model, diffusion = create_model_and_diffusion(29        **args_to_dict(args, model_and_diffusion_defaults().keys())30    )31    model.load_state_dict(32        dist_util.load_state_dict(args.model_path, map_location="cpu")33    )34    model.to(dist_util.dev())35    model.eval()36 37    logger.log("creating data loader...")38    data = load_data(39        data_dir=args.data_dir,40        batch_size=args.batch_size,41        image_size=args.image_size,42        class_cond=args.class_cond,43        deterministic=True,44    )45 46    logger.log("evaluating...")47    run_bpd_evaluation(model, diffusion, data, args.num_samples, args.clip_denoised)48 49 50def run_bpd_evaluation(model, diffusion, data, num_samples, clip_denoised):51    all_bpd = []52    all_metrics = {"vb": [], "mse": [], "xstart_mse": []}53    num_complete = 054    while num_complete < num_samples:55        batch, model_kwargs = next(data)56        batch = batch.to(dist_util.dev())57        model_kwargs = {k: v.to(dist_util.dev()) for k, v in model_kwargs.items()}58        minibatch_metrics = diffusion.calc_bpd_loop(59            model, batch, clip_denoised=clip_denoised, model_kwargs=model_kwargs60        )61 62        for key, term_list in all_metrics.items():63            terms = minibatch_metrics[key].mean(dim=0) / dist.get_world_size()64            dist.all_reduce(terms)65            term_list.append(terms.detach().cpu().numpy())66 67        total_bpd = minibatch_metrics["total_bpd"]68        total_bpd = total_bpd.mean() / dist.get_world_size()69        dist.all_reduce(total_bpd)70        all_bpd.append(total_bpd.item())71        num_complete += dist.get_world_size() * batch.shape[0]72 73        logger.log(f"done {num_complete} samples: bpd={np.mean(all_bpd)}")74 75    if dist.get_rank() == 0:76        for name, terms in all_metrics.items():77            out_path = os.path.join(logger.get_dir(), f"{name}_terms.npz")78            logger.log(f"saving {name} terms to {out_path}")79            np.savez(out_path, np.mean(np.stack(terms), axis=0))80 81    dist.barrier()82    logger.log("evaluation complete")83 84 85def create_argparser():86    defaults = dict(87        data_dir="", clip_denoised=True, num_samples=1000, batch_size=1, model_path=""88    )89    defaults.update(model_and_diffusion_defaults())90    parser = argparse.ArgumentParser()91    add_dict_to_argparser(parser, defaults)92    return parser93 94 95if __name__ == "__main__":96    main()97