multimodalart/pix2pix-zero
3
1import os, pdb2 3import argparse4import numpy as np5import torch6import requests7from PIL import Image8 9from diffusers import DDIMScheduler10from utils.ddim_inv import DDIMInversion11from utils.edit_directions import construct_direction12from utils.edit_pipeline import EditingPipeline13 14 15if __name__=="__main__":16 parser = argparse.ArgumentParser()17 parser.add_argument('--inversion', required=True)18 parser.add_argument('--prompt', type=str, required=True)19 parser.add_argument('--task_name', type=str, default='cat2dog')20 parser.add_argument('--results_folder', type=str, default='output/test_cat')21 parser.add_argument('--num_ddim_steps', type=int, default=50)22 parser.add_argument('--model_path', type=str, default='CompVis/stable-diffusion-v1-4')23 parser.add_argument('--xa_guidance', default=0.1, type=float)24 parser.add_argument('--negative_guidance_scale', default=5.0, type=float)25 parser.add_argument('--use_float_16', action='store_true')26 27 args = parser.parse_args()28 29 os.makedirs(os.path.join(args.results_folder, "edit"), exist_ok=True)30 os.makedirs(os.path.join(args.results_folder, "reconstruction"), exist_ok=True)31 32 if args.use_float_16:33 torch_dtype = torch.float1634 else:35 torch_dtype = torch.float3236 37 # if the inversion is a folder, the prompt should also be a folder38 assert (os.path.isdir(args.inversion)==os.path.isdir(args.prompt)), "If the inversion is a folder, the prompt should also be a folder"39 if os.path.isdir(args.inversion):40 l_inv_paths = sorted(glob(os.path.join(args.inversion, "*.pt")))41 l_bnames = [os.path.basename(x) for x in l_inv_paths]42 l_prompt_paths = [os.path.join(args.prompt, x.replace(".pt",".txt")) for x in l_bnames]43 else:44 l_inv_paths = [args.inversion]45 l_prompt_paths = [args.prompt]46 47 # Make the editing pipeline48 pipe = EditingPipeline.from_pretrained(args.model_path, torch_dtype=torch_dtype).to("cuda")49 pipe.scheduler = DDIMScheduler.from_config(pipe.scheduler.config)50 51 52 for inv_path, prompt_path in zip(l_inv_paths, l_prompt_paths):53 prompt_str = open(prompt_path).read().strip()54 rec_pil, edit_pil = pipe(prompt_str,55 num_inference_steps=args.num_ddim_steps,56 x_in=torch.load(inv_path).unsqueeze(0),57 edit_dir=construct_direction(args.task_name),58 guidance_amount=args.xa_guidance,59 guidance_scale=args.negative_guidance_scale,60 negative_prompt=prompt_str # use the unedited prompt for the negative prompt61 )62 63 bname = os.path.basename(args.inversion).split(".")[0]64 edit_pil[0].save(os.path.join(args.results_folder, f"edit/{bname}.png"))65 rec_pil[0].save(os.path.join(args.results_folder, f"reconstruction/{bname}.png"))66 