Team Ai
Modelpublic

MindoffAlex/HF-Diffusers-Deconstruct-Core99

sourceHugging Faceupdated 5mo agoView on Hugging Face
0likes
inversion_implemention.py206 linesDownload Raw Back to root
1import numpy as np2import PIL3from PIL import Image 4import torch5from torchvision import transforms6import diffusers7import torch8from transformers import CLIPTextModel, CLIPTokenizer9from diffusers import UNet2DConditionModel, AutoencoderKL, DDIMScheduler, DDIMInverseScheduler10from wandb import Image11 12 13class InversionImplementationDDIM:14    def __init__(self):15        16        self.model_name = "CompVis/stable-diffusion-v1-4"17        self.text_model_name = "openai/clip-vit-large-patch14"18        19        self.device = "cuda" if torch.cuda.is_available() else "cpu"20        self.dtype = torch.float16 if torch.cuda.is_available() else torch.float3221        22        self.tokenizer = CLIPTokenizer.from_pretrained(self.text_model_name)23        self.text_encoder = CLIPTextModel.from_pretrained(self.text_model_name)24        25        self.unet = UNet2DConditionModel.from_pretrained(self.model_name, subfolder="unet")26        self.vae = AutoencoderKL.from_pretrained(self.model_name, subfolder="vae")27        28        # 1. Standard scheduler for generation29        self.noise_scheduler = DDIMScheduler.from_pretrained(self.model_name, subfolder="scheduler")30        # 2. Inverse scheduler dedicated for inversion math31        self.inverse_scheduler = DDIMInverseScheduler.from_config(self.noise_scheduler.config)32        33        self.text_encoder.to(self.device, dtype=self.dtype).eval()34        self.unet.to(self.device, dtype=self.dtype).eval()35        self.vae.to(self.device, dtype=self.dtype).eval()36 37    38    @torch.inference_mode()39    def ddim_invers(self, num_inference_steps, init_image: str, prompt: str, visual_steps=None):40        if visual_steps is None:41            visual_steps = []42            43        # 1. Get the initial latent representation of the image (x_0)44        init_latents = self.get_latent_image(init_image)45 46        # 2. Get text embeddings for the prompt47        text_embeddings = self.get_text_embeddings(prompt)48 49        # 3. Use the inverse scheduler to set up timesteps going 0 -> 99950        self.inverse_scheduler.set_timesteps(num_inference_steps, device=self.device)51        timesteps = self.inverse_scheduler.timesteps 52 53        latents = init_latents.clone()54        saved_visuals = {}55 56        # 4. Perform DDIM inversion loop57        for idx, t in enumerate(timesteps):58            if idx in visual_steps:59                saved_visuals[idx] = self.vae_decoder(latents)60            # Predict the noise residual61            noise_pred = self.unet(latents, t, encoder_hidden_states=text_embeddings).sample62            63            # Step FORWARD in time (xt -> xt+1) using the specialized inverse scheduler64            # Note: DDIMInverseScheduler returns the next step in '.prev_sample'65            latents = self.inverse_scheduler.step(noise_pred, t, latents).prev_sample66            67        return latents, saved_visuals68 69 70    @torch.inference_mode()71    def ddim_sampling(self, num_inference_steps, inverted_latents, prompt: str, visual_steps=None):72        if visual_steps is None:73            visual_steps = []74            75        # 1. Setup standard scheduler timesteps (T -> 0)76        self.noise_scheduler.set_timesteps(num_inference_steps, device=self.device)77        timesteps = self.noise_scheduler.timesteps78 79        text_embeddings = self.get_text_embeddings(prompt)80        latents = inverted_latents.clone()81        saved_visuals = {}82 83        # 2. Standard generation loop84        for idx, t in enumerate(timesteps):85            if idx in visual_steps:86                saved_visuals[idx] = self.vae_decoder(latents)87            noise_pred = self.unet(latents, t, encoder_hidden_states=text_embeddings).sample88            # Steps backward in time to remove noise89            latents = self.noise_scheduler.step(noise_pred, t, latents).prev_sample90 91        # 3. Decode latents back to a PIL Image92        image = self.vae_decoder(latents)93        return image, saved_visuals94 95 96 97    @torch.inference_mode()98    def get_latent_image(self, image_path):99        image = PIL.Image.open(image_path).convert("RGB")100        preprocess = transforms.Compose([101            transforms.Resize((512, 512)),102            transforms.ToTensor(),103            transforms.Normalize([0.5], [0.5])104        ])105        image_tensor = (preprocess(image).unsqueeze(0).to(self.device, dtype=self.dtype))106 107        init_latents = self.vae.encode(image_tensor).latent_dist.mode()108        init_latents = init_latents * self.vae.config.scaling_factor109        110        return init_latents111 112 113    @torch.inference_mode()114    def get_text_embeddings(self, prompt):115        inputs = self.tokenizer(116            prompt,117            return_tensors="pt",118            padding="max_length",119            max_length=self.tokenizer.model_max_length,120            truncation=True,121        ).to(self.device)122 123        # 2. Extract the text embeddings124        text_embeddings = self.text_encoder(**inputs).last_hidden_state125 126        return text_embeddings127    128    129    @torch.inference_mode()130    def vae_decoder(self, latents):131        latents = 1 / 0.18215 * latents132        image = self.vae.decode(latents).sample133        image = (image / 2 + 0.5).clamp(0, 1)134        image = image.cpu().permute(0, 2, 3, 1).float().numpy()135        image = PIL.Image.fromarray((image[0] * 255).astype("uint8"))136        137        return image138 139 140 141def create_imageGrid(image_dict, output_filename):142    """Combines a dictionary of PIL images side-by-side into a single grid image."""143    if not image_dict:144        return145    146    # Sort by step index to ensure order147    sorted_steps = sorted(image_dict.keys())148    images = [image_dict[step] for step in sorted_steps]149    150    #extract size from the first PIL image object in the list151    width, height = images[0].size152    grid_width = width * len(images)153    grid_height = height154    155    #use PIL.Image.new to bypass wandb namespace conflict156    grid_img = PIL.Image.new("RGB", (grid_width, grid_height))157    158    # Paste images side by side159    for idx, img in enumerate(images):160        grid_img.paste(img, (idx * width, 0))161        162    grid_img.save(output_filename)163    print(f" Saved step grid to: {output_filename} (Steps: {sorted_steps})")164 165 166 167 168 169def main():170    steps = 50171    input_image_path = "/home/aviad/HF-Diffusers-Deconstruct-Core99/Road_in_Norway.jpg"  # Change to your file name172    prompt = "a photo of a road in norway"  # Describe your input image173    174    # Python indices: 0 = step 1, 1 = step 2, 2 = step 3175    steps_to_visualize = [0, 1, 2]176    177    print("Initializing DDIM Inversion System...")178    pipeline = InversionImplementationDDIM()179    180    # 1. Inversion181    print(f"Starting inversion...")182    inverted_noise, inversion_visuals = pipeline.ddim_invers(183        num_inference_steps=steps, 184        init_image=input_image_path, 185        prompt=prompt,186        visual_steps=steps_to_visualize187    )188    create_imageGrid(inversion_visuals, "grid_inversion_steps.jpg")189    190    # 2. Sampling191    print("Starting generation loop...")192    reconstructed_image, sampling_visuals = pipeline.ddim_sampling(193        num_inference_steps=steps, 194        inverted_latents=inverted_noise, 195        prompt=prompt,196        visual_steps=steps_to_visualize197    )198    create_imageGrid(sampling_visuals, "grid_sampling_steps.jpg")199    200    # Save the absolute final result201    reconstructed_image.save("reconstructed_final.jpg")202    print("Process complete! Check your folder for the grid files.")203 204 205if __name__ == "__main__":206    main()