CoolFace
Apppublic

Anonymous-123/ImageNet-Editing

sourceHugging Facecreativeml-openrail-mupdated 4y agoView on Hugging Face
1likes
super_res_sample.py120 linesDownload Raw Back to scripts
1"""2Generate a large batch of samples from a super resolution model, given a batch3of samples from a regular model from image_sample.py.4"""5 6import argparse7import os8 9import blobfile as bf10import numpy as np11import torch as th12import torch.distributed as dist13 14from guided_diffusion import dist_util, logger15from guided_diffusion.script_util import (16    sr_model_and_diffusion_defaults,17    sr_create_model_and_diffusion,18    args_to_dict,19    add_dict_to_argparser,20)21 22 23def main():24    args = create_argparser().parse_args()25 26    dist_util.setup_dist()27    logger.configure()28 29    logger.log("creating model...")30    model, diffusion = sr_create_model_and_diffusion(31        **args_to_dict(args, sr_model_and_diffusion_defaults().keys())32    )33    model.load_state_dict(34        dist_util.load_state_dict(args.model_path, map_location="cpu")35    )36    model.to(dist_util.dev())37    if args.use_fp16:38        model.convert_to_fp16()39    model.eval()40 41    logger.log("loading data...")42    data = load_data_for_worker(args.base_samples, args.batch_size, args.class_cond)43 44    logger.log("creating samples...")45    all_images = []46    while len(all_images) * args.batch_size < args.num_samples:47        model_kwargs = next(data)48        model_kwargs = {k: v.to(dist_util.dev()) for k, v in model_kwargs.items()}49        sample = diffusion.p_sample_loop(50            model,51            (args.batch_size, 3, args.large_size, args.large_size),52            clip_denoised=args.clip_denoised,53            model_kwargs=model_kwargs,54        )55        sample = ((sample + 1) * 127.5).clamp(0, 255).to(th.uint8)56        sample = sample.permute(0, 2, 3, 1)57        sample = sample.contiguous()58 59        all_samples = [th.zeros_like(sample) for _ in range(dist.get_world_size())]60        dist.all_gather(all_samples, sample)  # gather not supported with NCCL61        for sample in all_samples:62            all_images.append(sample.cpu().numpy())63        logger.log(f"created {len(all_images) * args.batch_size} samples")64 65    arr = np.concatenate(all_images, axis=0)66    arr = arr[: args.num_samples]67    if dist.get_rank() == 0:68        shape_str = "x".join([str(x) for x in arr.shape])69        out_path = os.path.join(logger.get_dir(), f"samples_{shape_str}.npz")70        logger.log(f"saving to {out_path}")71        np.savez(out_path, arr)72 73    dist.barrier()74    logger.log("sampling complete")75 76 77def load_data_for_worker(base_samples, batch_size, class_cond):78    with bf.BlobFile(base_samples, "rb") as f:79        obj = np.load(f)80        image_arr = obj["arr_0"]81        if class_cond:82            label_arr = obj["arr_1"]83    rank = dist.get_rank()84    num_ranks = dist.get_world_size()85    buffer = []86    label_buffer = []87    while True:88        for i in range(rank, len(image_arr), num_ranks):89            buffer.append(image_arr[i])90            if class_cond:91                label_buffer.append(label_arr[i])92            if len(buffer) == batch_size:93                batch = th.from_numpy(np.stack(buffer)).float()94                batch = batch / 127.5 - 1.095                batch = batch.permute(0, 3, 1, 2)96                res = dict(low_res=batch)97                if class_cond:98                    res["y"] = th.from_numpy(np.stack(label_buffer))99                yield res100                buffer, label_buffer = [], []101 102 103def create_argparser():104    defaults = dict(105        clip_denoised=True,106        num_samples=10000,107        batch_size=16,108        use_ddim=False,109        base_samples="",110        model_path="",111    )112    defaults.update(sr_model_and_diffusion_defaults())113    parser = argparse.ArgumentParser()114    add_dict_to_argparser(parser, defaults)115    return parser116 117 118if __name__ == "__main__":119    main()120