Team Ai
Modelpublic

MindoffAlex/HF-Diffusers-Deconstruct-Core99

sourceHugging Faceupdated 5mo agoView on Hugging Face
0likes
blended_loop.py237 linesDownload Raw Back to root
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()