Spaces:
Running on Zero
Running on Zero
Download utils.py from lea97338/LouPlayground: direct link, hf CLI and curl.
- Browser
- Download file 4.77 kB
-
https://huggingface.co/spaces/lea97338/LouPlayground/resolve/main/utils.py
- Command line
-
hf download hf://spaces/lea97338/LouPlayground/utils.py
-
curl -L -o utils.py https://huggingface.co/spaces/lea97338/LouPlayground/resolve/main/utils.py
4.77 kB
| 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 | |
| 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) |