Team Ai
Apppublic

TwoPerCent/instruct-pix2pix

sourceHugging Faceupdated 4y agoView on Hugging Face
0likes
edit_cli.py129 linesDownload Raw Back to root
1from __future__ import annotations2 3import math4import random5import sys6from argparse import ArgumentParser7 8import einops9import k_diffusion as K10import numpy as np11import torch12import torch.nn as nn13from einops import rearrange14from omegaconf import OmegaConf15from PIL import Image, ImageOps16from torch import autocast17 18sys.path.append("./stable_diffusion")19 20from stable_diffusion.ldm.util import instantiate_from_config21 22 23class CFGDenoiser(nn.Module):24    def __init__(self, model):25        super().__init__()26        self.inner_model = model27 28    def forward(self, z, sigma, cond, uncond, text_cfg_scale, image_cfg_scale):29        cfg_z = einops.repeat(z, "1 ... -> n ...", n=3)30        cfg_sigma = einops.repeat(sigma, "1 ... -> n ...", n=3)31        cfg_cond = {32            "c_crossattn": [torch.cat([cond["c_crossattn"][0], uncond["c_crossattn"][0], uncond["c_crossattn"][0]])],33            "c_concat": [torch.cat([cond["c_concat"][0], cond["c_concat"][0], uncond["c_concat"][0]])],34        }35        out_cond, out_img_cond, out_uncond = self.inner_model(cfg_z, cfg_sigma, cond=cfg_cond).chunk(3)36        return out_uncond + text_cfg_scale * (out_cond - out_img_cond) + image_cfg_scale * (out_img_cond - out_uncond)37 38 39def load_model_from_config(config, ckpt, vae_ckpt=None, verbose=False):40    print(f"Loading model from {ckpt}")41    pl_sd = torch.load(ckpt, map_location="cpu")42    if "global_step" in pl_sd:43        print(f"Global Step: {pl_sd['global_step']}")44    sd = pl_sd["state_dict"]45    if vae_ckpt is not None:46        print(f"Loading VAE from {vae_ckpt}")47        vae_sd = torch.load(vae_ckpt, map_location="cpu")["state_dict"]48        sd = {49            k: vae_sd[k[len("first_stage_model.") :]] if k.startswith("first_stage_model.") else v50            for k, v in sd.items()51        }52    model = instantiate_from_config(config.model)53    m, u = model.load_state_dict(sd, strict=False)54    if len(m) > 0 and verbose:55        print("missing keys:")56        print(m)57    if len(u) > 0 and verbose:58        print("unexpected keys:")59        print(u)60    return model61 62 63def main():64    parser = ArgumentParser()65    parser.add_argument("--resolution", default=512, type=int)66    parser.add_argument("--steps", default=100, type=int)67    parser.add_argument("--config", default="configs/generate.yaml", type=str)68    parser.add_argument("--ckpt", default="checkpoints/instruct-pix2pix-00-22000.ckpt", type=str)69    parser.add_argument("--vae-ckpt", default=None, type=str)70    parser.add_argument("--input", required=True, type=str)71    parser.add_argument("--output", required=True, type=str)72    parser.add_argument("--edit", required=True, type=str)73    parser.add_argument("--cfg-text", default=7.5, type=float)74    parser.add_argument("--cfg-image", default=1.5, type=float)75    parser.add_argument("--seed", type=int)76    args = parser.parse_args()77 78    config = OmegaConf.load(args.config)79    model = load_model_from_config(config, args.ckpt, args.vae_ckpt)80    model.eval().cuda()81    model_wrap = K.external.CompVisDenoiser(model)82    model_wrap_cfg = CFGDenoiser(model_wrap)83    null_token = model.get_learned_conditioning([""])84 85    seed = random.randint(0, 100000) if args.seed is None else args.seed86    input_image = Image.open(args.input).convert("RGB")87    width, height = input_image.size88    factor = args.resolution / max(width, height)89    factor = math.ceil(min(width, height) * factor / 64) * 64 / min(width, height)90    width = int((width * factor) // 64) * 6491    height = int((height * factor) // 64) * 6492    input_image = ImageOps.fit(input_image, (width, height), method=Image.Resampling.LANCZOS)93 94    if args.edit == "":95        input_image.save(args.output)96        return97 98    with torch.no_grad(), autocast("cuda"), model.ema_scope():99        cond = {}100        cond["c_crossattn"] = [model.get_learned_conditioning([args.edit])]101        input_image = 2 * torch.tensor(np.array(input_image)).float() / 255 - 1102        input_image = rearrange(input_image, "h w c -> 1 c h w").to(model.device)103        cond["c_concat"] = [model.encode_first_stage(input_image).mode()]104 105        uncond = {}106        uncond["c_crossattn"] = [null_token]107        uncond["c_concat"] = [torch.zeros_like(cond["c_concat"][0])]108 109        sigmas = model_wrap.get_sigmas(args.steps)110 111        extra_args = {112            "cond": cond,113            "uncond": uncond,114            "text_cfg_scale": args.cfg_text,115            "image_cfg_scale": args.cfg_image,116        }117        torch.manual_seed(seed)118        z = torch.randn_like(cond["c_concat"][0]) * sigmas[0]119        z = K.sampling.sample_euler_ancestral(model_wrap_cfg, z, sigmas, extra_args=extra_args)120        x = model.decode_first_stage(z)121        x = torch.clamp((x + 1.0) / 2.0, min=0.0, max=1.0)122        x = 255.0 * rearrange(x, "1 c h w -> h w c")123        edited_image = Image.fromarray(x.type(torch.uint8).cpu().numpy())124    edited_image.save(args.output)125 126 127if __name__ == "__main__":128    main()129