import torch import numpy as np from PIL import Image import torch.nn.functional as F class LouDiffusion: """Pipeline de diffusion DDPM optimisée pour Lou""" def __init__(self, model, device='cuda', num_train_timesteps=1000): self.model = model self.device = device self.num_train_timesteps = num_train_timesteps # Beta schedule - cosine pour génération rapide self.betas = self._cosine_beta_schedule(num_train_timesteps).to(device) self.alphas = 1.0 - self.betas self.alphas_cumprod = torch.cumprod(self.alphas, dim=0) self.alphas_cumprod_prev = F.pad(self.alphas_cumprod[:-1], (1, 0), value=1.0) self.sqrt_alphas_cumprod = torch.sqrt(self.alphas_cumprod) self.sqrt_one_minus_alphas_cumprod = torch.sqrt(1.0 - self.alphas_cumprod) def _cosine_beta_schedule(self, timesteps, s=0.008): steps = timesteps + 1 x = torch.linspace(0, timesteps, steps) alphas_cumprod = torch.cos(((x / timesteps) + s) / (1 + s) * torch.pi * 0.5) ** 2 alphas_cumprod = alphas_cumprod / alphas_cumprod[0] betas = 1 - (alphas_cumprod[1:] / alphas_cumprod[:-1]) return torch.clip(betas, 0.0001, 0.9999) def add_noise(self, x_start, t, noise=None): if noise is None: noise = torch.randn_like(x_start) sqrt_alphas_cumprod_t = self.sqrt_alphas_cumprod[t][:, None, None, None] sqrt_one_minus_alphas_cumprod_t = self.sqrt_one_minus_alphas_cumprod[t][:, None, None, None] return sqrt_alphas_cumprod_t * x_start + sqrt_one_minus_alphas_cumprod_t * noise @torch.no_grad() def sample(self, shape, num_inference_steps=50, guidance_scale=7.5, prompt_embeds=None, generator=None): """Sampling avec classifier-free guidance""" self.model.eval() batch_size = shape[0] # Scheduler pour sampling rapide timesteps = torch.linspace(self.num_train_timesteps - 1, 0, num_inference_steps, dtype=torch.long, device=self.device) # Bruit initial x = torch.randn(shape, device=self.device, generator=generator) for i, t in enumerate(timesteps): t_batch = torch.full((batch_size,), t, device=self.device, dtype=torch.long) # Prédiction conditionnée noise_pred = self.model(x, t_batch) # Guidance scale (simplifié) if guidance_scale != 1.0: noise_pred_uncond = self.model(x, t_batch) noise_pred = noise_pred_uncond + guidance_scale * (noise_pred - noise_pred_uncond) # DDPM step alpha_t = self.alphas_cumprod[t] alpha_t_prev = self.alphas_cumprod_prev[t] if t > 0 else torch.tensor(1.0, device=self.device) pred_x0 = (x - torch.sqrt(1 - alpha_t) * noise_pred) / torch.sqrt(alpha_t) pred_x0 = torch.clamp(pred_x0, -1, 1) direction = torch.sqrt(1 - alpha_t_prev) * noise_pred x = torch.sqrt(alpha_t_prev) * pred_x0 + direction if i < len(timesteps) - 1: noise = torch.randn_like(x) if i > 0 else torch.zeros_like(x) x = x + torch.sqrt(self.betas[t]) * noise return torch.clamp(x, -1, 1) def training_step(self, images): """Un pas d'entraînement""" self.model.train() batch_size = images.shape[0] device = images.device # Timesteps aléatoires t = torch.randint(0, self.num_train_timesteps, (batch_size,), device=device) # Bruit noise = torch.randn_like(images) # Images bruitées noisy_images = self.add_noise(images, t, noise) # Prédiction noise_pred = self.model(noisy_images, t) # Loss MSE loss = F.mse_loss(noise_pred, noise) return loss def tensor_to_pil(tensor): """Convertit un tensor [-1, 1] en PIL Image""" tensor = (tensor + 1) / 2 tensor = torch.clamp(tensor, 0, 1) tensor = tensor.cpu().permute(0, 2, 3, 1).numpy() images = (tensor * 255).astype(np.uint8) return [Image.fromarray(img) for img in images] def pil_to_tensor(images, size=256): """Convertit PIL Images en tensor [-1, 1]""" if not isinstance(images, list): images = [images] tensors = [] for img in images: img = img.convert('RGB').resize((size, size)) arr = np.array(img).astype(np.float32) / 255.0 tensor = torch.from_numpy(arr).permute(2, 0, 1) tensor = tensor * 2 - 1 tensors.append(tensor) return torch.stack(tensors)