MindoffAlex/HF-Diffusers-Deconstruct-Core99
0
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() 