LouPlayground / utils.py
lea97338's picture
Upload 4 files
4af0d61 verified
Raw History Blame Contribute Delete
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
@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)