import torch
import torch.nn as nn
import torch.nn.functional as F
import numpy as np
from typing import Tuple
class VectorQuantizer(nn.Module):
"""Quantification vectorielle avec straight-through estimator.
Implémente la quantification par plus proche voisin dans le codebook,
avec passage straight-through du gradient pour permettre la
rétropropagation à travers l'opération d'indexation discrète.
Args:
n_embeddings: Nombre d'entrées dans le codebook (taille K).
embedding_dim: Dimension des vecteurs d'embedding.
commitment_cost: Poids de la perte de commitment (β dans l'article).
"""
def __init__(self, n_embeddings: int = 1024,
embedding_dim: int = 256,
commitment_cost: float = 0.25):
super().__init__()
self.n_embeddings = n_embeddings
self.embedding_dim = embedding_dim
self.commitment_cost = commitment_cost
self.embedding = nn.Embedding(n_embeddings, embedding_dim)
self.embedding.weight.data.uniform_(
-1.0 / n_embeddings, 1.0 / n_embeddings
)
def forward(self, z: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor,
torch.Tensor, torch.Tensor]:
"""Quantifie les vecteurs latents z.
Args:
z: Tenseur d'entrée (B, D, H, W).
Returns:
Tuple (z_q, indices, perte_commitment, perplexité).
"""
z_flat = z.permute(0, 2, 3, 1).reshape(-1, self.embedding_dim)
distances = torch.cdist(z_flat, self.embedding.weight)
indices = distances.argmin(dim=-1)
z_q = self.embedding(indices).view(z.shape)
commitment_loss = F.mse_loss(z_q.detach(), z) * self.commitment_cost
z_q = z + (z_q - z).detach()
encodings = F.one_hot(indices, self.n_embeddings).float()
avg_probs = encodings.mean(dim=0)
perplexity = torch.exp(-torch.sum(
avg_probs * torch.log(avg_probs + 1e-10)
))
return z_q, indices, commitment_loss, perplexity
class DiscriminateurPatchGAN(nn.Module):
"""Discriminateur PatchGAN 70×70.
Classifie chaque patch 70×70 de l'image comme réel ou faux,
forçant le décodeur à produire des détails haute-fréquence locaux.
"""
def __init__(self, in_channels: int = 3):
super().__init__()
self.layers = nn.Sequential(
nn.Conv2d(in_channels, 64, kernel_size=4, stride=2, padding=1),
nn.LeakyReLU(0.2, inplace=True),
nn.Conv2d(64, 128, kernel_size=4, stride=2, padding=1),
nn.BatchNorm2d(128),
nn.LeakyReLU(0.2, inplace=True),
nn.Conv2d(128, 256, kernel_size=4, stride=2, padding=1),
nn.BatchNorm2d(256),
nn.LeakyReLU(0.2, inplace=True),
nn.Conv2d(256, 1, kernel_size=4, stride=1, padding=1),
)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.layers(x)
class VQGAN(nn.Module):
"""Modèle VQGAN complet avec encodeur, quantizer et décodeur.
Args:
in_channels: Nombre de canaux d'entrée (3 pour RGB).
latent_dim: Dimension de l'espace latent.
n_embeddings: Taille du codebook.
commitment_cost: Poids de la perte de commitment.
"""
def __init__(self, in_channels: int = 3, latent_dim: int = 256,
n_embeddings: int = 1024, commitment_cost: float = 0.25):
super().__init__()
self.latent_dim = latent_dim
self.encoder = nn.Sequential(
nn.Conv2d(in_channels, 128, kernel_size=4, stride=2, padding=1),
nn.BatchNorm2d(128),
nn.ReLU(inplace=True),
nn.Conv2d(128, 256, kernel_size=4, stride=2, padding=1),
nn.BatchNorm2d(256),
nn.ReLU(inplace=True),
nn.Conv2d(256, latent_dim, kernel_size=3, stride=1, padding=1),
nn.BatchNorm2d(latent_dim),
)
self.quantizer = VectorQuantizer(
n_embeddings, latent_dim, commitment_cost
)
self.decoder = nn.Sequential(
nn.ConvTranspose2d(latent_dim, 256,
kernel_size=3, stride=1, padding=1),
nn.BatchNorm2d(256),
nn.ReLU(inplace=True),
nn.ConvTranspose2d(256, 128,
kernel_size=4, stride=2, padding=1),
nn.BatchNorm2d(128),
nn.ReLU(inplace=True),
nn.ConvTranspose2d(128, in_channels,
kernel_size=4, stride=2, padding=1),
nn.Tanh(),
)
def forward(self, x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor,
torch.Tensor, torch.Tensor]:
"""Passe avant complète : encodage → quantification → décodage.
Args:
x: Image d'entrée (B, 3, H, W), normalisée dans [-1, 1].
Returns:
Tuple (x_reconstruit, indices, perte_commitment, perplexité).
"""
z = self.encoder(x)
z_q, indices, commitment_loss, perplexity = self.quantizer(z)
x_hat = self.decoder(z_q)
return x_hat, indices, commitment_loss, perplexity
def encode_to_indices(self, x: torch.Tensor) -> torch.Tensor:
"""Encode une image en indices du codebook.
Args:
x: Image d'entrée (B, 3, H, W).
Returns:
Indices du codebook (B, H/4 * W/4).
"""
z = self.encoder(x)
_, indices, _, _ = self.quantizer(z)
return indices
def perte_vqgan(model: VQGAN, discriminateur: DiscriminateurPatchGAN,
x: torch.Tensor, lambda_adv: float = 0.1,
lambda_percept: float = 1.0) -> Tuple[torch.Tensor, dict]:
"""Calcule la perte combinée pour l'entraînement VQGAN.
Args:
model: Modèle VQGAN.
discriminateur: Discriminateur PatchGAN.
x: Images réelles (B, 3, H, W).
lambda_adv: Poids de la perte adverse.
lambda_percept: Poids de la perte perceptuelle.
Returns:
Tuple (perte_totale, dictionnaire des composantes de perte).
"""
x_hat, indices, commitment_loss, perplexity = model(x)
rec_loss_l1 = F.l1_loss(x_hat, x)
rec_loss_l2 = F.mse_loss(x_hat, x)
logits_fake = discriminateur(x_hat)
adv_loss = F.binary_cross_entropy_with_logits(
logits_fake, torch.ones_like(logits_fake)
)
total_loss = (rec_loss_l1 + rec_loss_l2
+ commitment_loss
+ lambda_adv * adv_loss)
métriques = {
'l1': rec_loss_l1.item(),
'l2': rec_loss_l2.item(),
'commitment': commitment_loss.item(),
'adversarial': adv_loss.item(),
'perplexité': perplexity.item(),
'utilisation_codebook': perplexity.item() / model.quantizer.n_embeddings,
}
return total_loss, métriques
import lpips
def évaluer_compression(VQGAN, dataloader: torch.utils.data.DataLoader,
device: str = 'cuda') -> dict:
"""Évalue la qualité de compression VQGAN sur un ensemble de test.
Args:
model: Modèle VQGAN entraîné.
dataloader: DataLoader d'images de test.
device: Périphérique de calcul.
Returns:
dict: Métriques moyennes (PSNR, SSIM, LPIPS, bitrate, FID si calculable).
"""
model = VQGAN.to(device).eval()
lpips_fn = lpips.LPIPS(net='alex').to(device)
psnr_total = 0.0
lpips_total = 0.0
bitrates = []
n = 0
with torch.no_grad():
for x, _ in dataloader:
x = x.to(device)
x_hat, indices, _, _ = model(x)
mse = F.mse_loss(x_hat, x).item()
psnr = 10 * np.log10(4.0 / mse)
psnr_total += psnr
lpips_val = lpips_fn(x, x_hat).mean().item()
lpips_total += lpips_val
for i in range(x.shape[0]):
bpp = calculer_bitrate(model, indices[i:i+1],
x.shape[2], x.shape[3])
bitrates.append(bpp)
n += x.shape[0]
return {
'PSNR (dB)': psnr_total / n,
'LPIPS': lpips_total / n,
'Bitrate moyen (bpp)': np.mean(bitrates),
'Bitrate min (bpp)': np.min(bitrates),
'Bitrate max (bpp)': np.max(bitrates),
}
def entraîner_vqgan(model: VQGAN, dataloader, n_epochs: int = 100,
lr: float = 1e-4, device: str = 'cuda'):
"""Boucle d'entraînement VQGAN complète.
Args:
model: Modèle VQGAN à entraîner.
dataloader: DataLoader d'entraînement.
n_epochs: Nombre d'époques.
lr: Taux d'apprentissage.
device: Périphérique de calcul.
Note:
Surveiller la perplexité du codebook : si elle descend
en dessous de 10% de K (taille du codebook), activer l'EMA
update ou réinitialiser les embeddings inutilisés.
"""
model = model.to(device)
disc = DiscriminateurPatchGAN().to(device)
optim_g = torch.optim.Adam(model.parameters(), lr=lr)
optim_d = torch.optim.Adam(disc.parameters(), lr=lr * 0.5)
for epoch in range(n_epochs):
for batch_idx, (x, _) in enumerate(dataloader):
x = x.to(device)
optim_g.zero_grad()
loss_g, métriques = perte_vqgan(model, disc, x)
loss_g.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
optim_g.step()
optim_d.zero_grad()
with torch.no_grad():
x_hat, _, _, _ = model(x)
logits_real = disc(x)
logits_fake = disc(x_hat)
loss_d_real = F.binary_cross_entropy_with_logits(
logits_real, torch.ones_like(logits_real)
)
loss_d_fake = F.binary_cross_entropy_with_logits(
logits_fake, torch.zeros_like(logits_fake)
)
loss_d = (loss_d_real + loss_d_fake) * 0.5
loss_d.backward()
optim_d.step()
if batch_idx % 100 == 0:
print(f"Epoch {epoch:3d} | Batch {batch_idx:4d} | "
f"G: {loss_g.item():.3f} | D: {loss_d.item():.3f} | "
f"Perplex: {métriques['perplexité']:.1f} | "
f"Codebook: {métriques['utilisation_codebook']:.1%}")
torch.save({
'state_dict': model.state_dict(),
'disc_state_dict': disc.state_dict(),
'optim_g': optim_g.state_dict(),
'optim_d': optim_d.state_dict(),
}, 'vqgan_entraîné.pt')