MindoffAlex/HF-Diffusers-Deconstruct-Core99
0
1import PIL.Image2import diffusers3import numpy as np4import diffusers5from transformers import CLIPTextModel, CLIPTokenizer6from torchvision import transforms7import torch 8 9class BlendedLatentDiffusion:10 def __init__(self):11 model_name = "CompVis/stable-diffusion-v1-4"12 text_model_name = "openai/clip-vit-large-patch14"13 14 # 1. Hardware configuration15 self.device = "cuda" if torch.cuda.is_available() else "cpu"16 # Use float16 for speed and lower VRAM if running on GPU, otherwise float3217 self.dtype = torch.float16 if self.device == "cuda" else torch.float3218 19 # 2. Load all model components20 self.autoencoder = diffusers.AutoencoderKL.from_pretrained(model_name, subfolder="vae")21 self.text_encoder = CLIPTextModel.from_pretrained(text_model_name)22 self.tokenizer = CLIPTokenizer.from_pretrained(text_model_name)23 self.unet = diffusers.UNet2DConditionModel.from_pretrained(model_name, subfolder="unet")24 25 # 3. Load the Scheduler (DDIMScheduler is ideal for Blended Diffusion)26 # self.scheduler = diffusers.DDIMScheduler.from_pretrained(model_name, subfolder="scheduler")27 self.scheduler = diffusers.DPMSolverMultistepScheduler.from_pretrained(model_name, subfolder="scheduler", algorithm_type="dpmsolver++", use_karras_sigmas=True) # Alternative scheduler for experimentation28 29 30 # 4. Cast and move models to target device31 self.autoencoder.to(device=self.device, dtype=self.dtype).eval()32 self.text_encoder.to(device=self.device, dtype=self.dtype).eval()33 self.unet.to(device=self.device, dtype=self.dtype).eval()34 35 36 37 def blended_latent_diffusion(self,38 init_image: PIL.Image, 39 mask_image: PIL.Image, 40 prompt: str, 41 num_inference_steps: int = 25,42 strength: float = 0.8,43 guidance_scale: float = 7.544 ) -> PIL.Image:45 """46 Applies blended latent diffusion to an input image using a mask and a text prompt.47 48 Args:49 init_image (PIL.Image): The initial image to be modified.50 mask_image (PIL.Image): A binary mask image where white areas indicate regions to modify.51 prompt (str): The text prompt guiding the diffusion process.52 num_inference_steps (int): The number of inference steps for the diffusion process.53 strength (float): The strength of the diffusion effect, between 0 and 1.54 guidance_scale (float): The scale for guidance, controlling the influence of the text prompt.55 Returns:56 PIL.Image: The modified image after applying blended latent diffusion.57 """58 59 print(f"🚀 Starting Blended Latent Diffusion | Prompt: '{prompt}'")60 print(f"📦 Configurations | Steps: {num_inference_steps} | CFG Scale: {guidance_scale}")61 62 # Step 1: Preprocess the input images63 init_image = init_image.convert("RGB")64 mask_image = mask_image.convert("L") # Convert to grayscale for masking65 66 print("⏳ Encoding initial image to latent space...")67 # Step 2: Encode the initial image into latent space and transform the mask68 latent_init = self.encode_to_latent(init_image)69 mask_transform = self.preprocess_mask(mask_image)70 print(f"✅ Latents Prepared | Shape: {list(latent_init.shape)} | Mask Shape: {list(mask_transform.shape)}")71 72 # Step 3: Generate noise based on the prompt73 print("⏳ Processing text prompt and creating base noise...")74 noise = self.generate_noise_from_prompt(prompt, latent_init.shape)75 print(f"✅ Text Embeddings Configured | Embedded Shape: {list(noise[1].shape)}")76 77 # Step 4: Blend the noise with the latent representation using the mask78 print(f"⏳ Entering Denoising Loop ({num_inference_steps} steps via {self.scheduler.__class__.__name__})...")79 blended_latent = self.blend_latent_with_mask(80 latent_init, noise, mask_transform, strength, num_inference_steps, guidance_scale)81 print("✅ Latent optimization sequence complete.")82 83 print("⏳ Decoding final blended latents back to image pixels...")84 # Step 5: Decode the blended latent representation back to an image85 output_image = self.decode_from_latent(blended_latent)86 print("✨ Process complete! Returning output image.")87 88 return output_image89 90 91 def encode_to_latent(self, init_image: PIL.Image) -> torch.Tensor:92 preprocess = transforms.Compose([93 transforms.Resize((512, 512)), 94 transforms.ToTensor(), 95 transforms.Normalize([0.5], [0.5]) 96 ])97 input_tensor = preprocess(init_image).unsqueeze(0).to(self.device, dtype=self.dtype)98 with torch.no_grad():99 latents = self.autoencoder.encode(input_tensor).latent_dist.sample()100 return latents * 0.18215101 102 def preprocess_mask(self, mask_image: PIL.Image) -> torch.Tensor:103 # Resize to latent space size (512 / 8 = 64)104 mask = mask_image.resize((64, 64), resample=PIL.Image.NEAREST)105 mask = transforms.ToTensor()(mask).to(self.device, dtype=self.dtype) # Shape: [1, 64, 64]106 107 # FIX: Add a batch dimension to make it [1, 1, 64, 64] for clean matrix broadcasting108 return mask.unsqueeze(0) 109 110 111 112 def generate_noise_from_prompt(self, prompt: str, latent_shape: torch.Size) -> tuple[torch.Tensor, torch.Tensor]:113 """114 Prepares text context with CFG support and creates the baseline noise vector.115 """116 # 1. Encode the positive conditional prompt117 text_inputs = self.tokenizer(118 prompt, padding="max_length", max_length=self.tokenizer.model_max_length, return_tensors="pt"119 )120 text_embeddings = self.text_encoder(text_inputs.input_ids.to(self.device)).last_hidden_state121 122 # 2. Encode the unconditional empty prompt (negative guidance)123 uncond_inputs = self.tokenizer(124 "", padding="max_length", max_length=self.tokenizer.model_max_length, return_tensors="pt"125 )126 uncond_embeddings = self.text_encoder(uncond_inputs.input_ids.to(self.device)).last_hidden_state127 128 # 3. Concatenate them into a single batch for parallel UNet processing129 # Shape becomes [2, 77, 768]130 text_embeddings = torch.cat([uncond_embeddings, text_embeddings])131 132 # 4. Generate the single static base noise layout133 init_noise = torch.randn(latent_shape, device=self.device, dtype=self.dtype)134 135 return init_noise, text_embeddings136 137 138 def blend_latent_with_mask(139 self, 140 latent_init: torch.Tensor, 141 noise_package: tuple[torch.Tensor, torch.Tensor], 142 mask_tensor: torch.Tensor, 143 strength: float,144 num_inference_steps: int = 25,145 guidance_scale: float = 7.5 146 ) -> torch.Tensor:147 """148 Executes Blended Latent Diffusion using DPMSolverMultistepScheduler149 with strict 1D tensor array conversion for add_noise compatibility.150 """151 init_noise, text_embeddings = noise_package 152 153 # 1. Initialize full steps on the scheduler154 self.scheduler.set_timesteps(num_inference_steps, device=self.device)155 156 # 2. Slice timesteps based on strength parameter157 init_timestep_idx = int(num_inference_steps * (1 - strength))158 timesteps = self.scheduler.timesteps[init_timestep_idx:]159 160 # 3. Configure multi-step tracking properties161 if hasattr(self.scheduler, "set_begin_index"):162 self.scheduler.set_begin_index(init_timestep_idx)163 164 # 4. FIX: Force the starting step to be a 1D vector tensor to prevent IndexError165 start_t = timesteps[0].item()166 start_timestep_tensor = torch.tensor([start_t], device=self.device, dtype=torch.long)167 168 # Initialize foreground latents with properly scaled starting noise169 latents_fg = self.scheduler.add_noise(latent_init, init_noise, start_timestep_tensor)170 171 # 5. Generate a single background noise layout to maintain calculation history172 fresh_bg_noise = torch.randn_like(latent_init)173 174 # 6. Core Denoising Loop175 for idx, t in enumerate(timesteps):176 # A. FIX: Force the loop timestep 't' into a 1D vector tensor for add_noise safety177 current_t_val = t.item() if isinstance(t, torch.Tensor) else t178 t_tensor = torch.tensor([current_t_val], device=self.device, dtype=torch.long)179 180 # Prepare background for current timestep 't' using the 1D tensor181 latents_bg = self.scheduler.add_noise(latent_init, fresh_bg_noise, t_tensor)182 183 # B. Spatial Blending: Sync background state to keep boundaries crisp184 latents_fg = mask_tensor * latents_fg + (1.0 - mask_tensor) * latents_bg185 186 # C. Duplicate inputs for Classifier-Free Guidance (CFG) processing187 latent_model_input = torch.cat([latents_fg] * 2)188 latent_model_input = self.scheduler.scale_model_input(latent_model_input, t)189 190 # D. Predict noise maps using the UNet configuration191 with torch.no_grad():192 noise_pred = self.unet(latent_model_input, t, encoder_hidden_states=text_embeddings).sample193 194 # E. Split predictions and extrapolate prompt guidance strength195 noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)196 noise_pred = noise_pred_uncond + guidance_scale * (noise_pred_text - noise_pred_uncond)197 198 # F. Step foreground latents backward one step in time199 latents_fg = self.scheduler.step(noise_pred, t, latents_fg).prev_sample200 201 # 7. Final blending pass at t=0 to keep unmasked original pixels perfectly clean202 latents_fg = mask_tensor * latents_fg + (1.0 - mask_tensor) * latent_init203 204 return latents_fg205 206 207 def decode_from_latent(self, blended_latent: torch.Tensor) -> PIL.Image:208 # Undo the VAE scaling factor209 latents = blended_latent / 0.18215210 with torch.no_grad():211 image_tensor = self.autoencoder.decode(latents).sample212 213 # Convert tensor back to PIL Image214 image_tensor = (image_tensor / 2 + 0.5).clamp(0, 1) # Rescale back to [0, 1]215 image_tensor = image_tensor.cpu().permute(0, 2, 3, 1).float().numpy()216 image_numpy = (image_tensor * 255).astype("uint8")[0]217 return PIL.Image.fromarray(image_numpy)218 219 220 221def main():222 blended_diffusion = BlendedLatentDiffusion()223 224 init_image = PIL.Image.open("/home/aviad/interview/mobileye/messi.jpg")225 mask_image = PIL.Image.open("/home/aviad/interview/mobileye/messi_mask.png")226 output = blended_diffusion.blended_latent_diffusion(227 init_image=init_image,228 mask_image=mask_image,229 prompt="fluffy white clouds in a bright blue sky, highly detailed",230 num_inference_steps=25,231 strength=0.95, # High strength allows completely overwriting the target area232 guidance_scale=12.0 # Slightly higher scale forces strong prompt adhesion over background textures233 )234 output.save("output_image.jpg")235 236if __name__ == "__main__": 237 main()