Spaces:
Sleeping
Sleeping
| """ | |
| model_bce_bn.py — Conditional VAE para HAM10000 | |
| Versión híbrida: | |
| - Basada en la arquitectura visualmente buena: | |
| BatchNorm + ConvTranspose2d + BCE | |
| - Con fixes de estabilidad: | |
| logvar clamp, std clamp, z clamp | |
| - Usa BCEWithLogitsLoss: | |
| el decoder devuelve logits durante entrenamiento | |
| sigmoid se aplica solo al generar/reconstruir imágenes | |
| IMPORTANTE: | |
| - Las imágenes de entrada deben estar en [0, 1] | |
| - En train.py usar transforms.ToTensor() SIN Normalize() | |
| """ | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| # ─── Constantes ───────────────────────────────────────────────────────────── | |
| IMG_SIZE = 128 | |
| LATENT_DIM = 128 | |
| NUM_CLASSES = 7 | |
| # Ojo: ImageFolder ordena alfabéticamente si usas carpetas. | |
| # Para tu dataset actual normalmente será: | |
| # ["akiec", "bcc", "bkl", "df", "mel", "nv", "vasc"] | |
| CLASS_NAMES = ["akiec", "bcc", "bkl", "df", "mel", "nv", "vasc"] | |
| CLASS_DESCRIPTIONS = { | |
| "mel": "Melanoma — cáncer agresivo de melanocitos", | |
| "nv": "Nevi melanocíticos — lunares benignos", | |
| "bcc": "Carcinoma basocelular — cáncer frecuente, lento", | |
| "akiec": "Queratosis actínica / Bowen — lesión precancerosa", | |
| "bkl": "Queratosis benigna — manchas benignas", | |
| "df": "Dermatofibroma — bulto benigno firme", | |
| "vasc": "Lesiones vasculares — hemangiomas y similares", | |
| } | |
| # Estabilidad numérica | |
| LOGVAR_MIN = -6.0 | |
| LOGVAR_MAX = 2.0 | |
| STD_MAX = 3.0 | |
| Z_CLAMP = 10.0 | |
| # ─── Inicialización ───────────────────────────────────────────────────────── | |
| def init_weights(module): | |
| if isinstance(module, (nn.Conv2d, nn.ConvTranspose2d, nn.Linear)): | |
| nn.init.xavier_normal_(module.weight) | |
| if module.bias is not None: | |
| nn.init.zeros_(module.bias) | |
| elif isinstance(module, nn.BatchNorm2d): | |
| nn.init.ones_(module.weight) | |
| nn.init.zeros_(module.bias) | |
| # ─── Bloque residual ──────────────────────────────────────────────────────── | |
| class ResBlock(nn.Module): | |
| def __init__(self, channels): | |
| super().__init__() | |
| self.block = nn.Sequential( | |
| nn.Conv2d(channels, channels, kernel_size=3, padding=1, bias=False), | |
| nn.BatchNorm2d(channels), | |
| nn.ReLU(inplace=True), | |
| nn.Conv2d(channels, channels, kernel_size=3, padding=1, bias=False), | |
| nn.BatchNorm2d(channels), | |
| ) | |
| def forward(self, x): | |
| return F.relu(x + self.block(x), inplace=True) | |
| # ─── Encoder ──────────────────────────────────────────────────────────────── | |
| class Encoder(nn.Module): | |
| """ | |
| Encoder: | |
| 128x128x3 | |
| → 64x64x32 | |
| → 32x32x64 | |
| → 16x16x128 | |
| → 8x8x256 | |
| → mu, logvar | |
| """ | |
| def __init__(self, latent_dim=LATENT_DIM): | |
| super().__init__() | |
| self.conv = nn.Sequential( | |
| # 128x128x3 → 64x64x32 | |
| nn.Conv2d(3, 32, kernel_size=4, stride=2, padding=1, bias=False), | |
| nn.BatchNorm2d(32), | |
| nn.ReLU(inplace=True), | |
| ResBlock(32), | |
| # 64x64x32 → 32x32x64 | |
| nn.Conv2d(32, 64, kernel_size=4, stride=2, padding=1, bias=False), | |
| nn.BatchNorm2d(64), | |
| nn.ReLU(inplace=True), | |
| ResBlock(64), | |
| # 32x32x64 → 16x16x128 | |
| nn.Conv2d(64, 128, kernel_size=4, stride=2, padding=1, bias=False), | |
| nn.BatchNorm2d(128), | |
| nn.ReLU(inplace=True), | |
| ResBlock(128), | |
| # 16x16x128 → 8x8x256 | |
| nn.Conv2d(128, 256, kernel_size=4, stride=2, padding=1, bias=False), | |
| nn.BatchNorm2d(256), | |
| nn.ReLU(inplace=True), | |
| ) | |
| self.flatten_dim = 256 * 8 * 8 | |
| self.fc_mu = nn.Linear(self.flatten_dim, latent_dim) | |
| self.fc_logvar = nn.Linear(self.flatten_dim, latent_dim) | |
| self.apply(init_weights) | |
| # Arranque más estable: varianza inicial pequeña. | |
| nn.init.constant_(self.fc_logvar.bias, -2.0) | |
| def forward(self, x): | |
| h = self.conv(x) | |
| h = h.view(h.size(0), -1) | |
| mu = self.fc_mu(h) | |
| logvar = self.fc_logvar(h) | |
| logvar = torch.clamp(logvar, min=LOGVAR_MIN, max=LOGVAR_MAX) | |
| return mu, logvar | |
| # ─── Decoder ──────────────────────────────────────────────────────────────── | |
| class Decoder(nn.Module): | |
| """ | |
| Decoder: | |
| z + class_embedding | |
| → 8x8x256 | |
| → 16x16x128 | |
| → 32x32x64 | |
| → 64x64x32 | |
| → 128x128x3 logits | |
| IMPORTANTE: | |
| No usamos Sigmoid aquí durante training. | |
| La loss usa BCEWithLogitsLoss. | |
| """ | |
| def __init__(self, latent_dim=LATENT_DIM, num_classes=NUM_CLASSES): | |
| super().__init__() | |
| self.class_embed = nn.Embedding(num_classes, 32) | |
| self.fc = nn.Linear(latent_dim + 32, 256 * 8 * 8) | |
| self.deconv = nn.Sequential( | |
| ResBlock(256), | |
| # 8x8x256 → 16x16x128 | |
| nn.ConvTranspose2d(256, 128, kernel_size=4, stride=2, padding=1, bias=False), | |
| nn.BatchNorm2d(128), | |
| nn.ReLU(inplace=True), | |
| ResBlock(128), | |
| # 16x16x128 → 32x32x64 | |
| nn.ConvTranspose2d(128, 64, kernel_size=4, stride=2, padding=1, bias=False), | |
| nn.BatchNorm2d(64), | |
| nn.ReLU(inplace=True), | |
| ResBlock(64), | |
| # 32x32x64 → 64x64x32 | |
| nn.ConvTranspose2d(64, 32, kernel_size=4, stride=2, padding=1, bias=False), | |
| nn.BatchNorm2d(32), | |
| nn.ReLU(inplace=True), | |
| ResBlock(32), | |
| # 64x64x32 → 128x128x3 logits | |
| nn.ConvTranspose2d(32, 3, kernel_size=4, stride=2, padding=1), | |
| ) | |
| self.apply(init_weights) | |
| def forward(self, z, class_label): | |
| ce = self.class_embed(class_label) | |
| zc = torch.cat([z, ce], dim=1) | |
| h = self.fc(zc) | |
| h = h.view(h.size(0), 256, 8, 8) | |
| logits = self.deconv(h) | |
| return logits | |
| # ─── CVAE completo ────────────────────────────────────────────────────────── | |
| class ConditionalVAE(nn.Module): | |
| """ | |
| Conditional Variational Autoencoder para HAM10000. | |
| Entrada: | |
| x en [0, 1] | |
| Forward: | |
| devuelve x_logits, mu, logvar | |
| Loss: | |
| BCEWithLogits(x_logits, x) + beta * KL | |
| """ | |
| def __init__(self, latent_dim=LATENT_DIM, num_classes=NUM_CLASSES, beta=1.0): | |
| super().__init__() | |
| self.latent_dim = latent_dim | |
| self.beta = beta | |
| self.encoder = Encoder(latent_dim) | |
| self.decoder = Decoder(latent_dim, num_classes) | |
| def reparametrize(self, mu, logvar, sample=True): | |
| """ | |
| z = mu + eps * std | |
| sample=True: | |
| usado en training. | |
| sample=False: | |
| usado en validación/reconstrucción determinista. | |
| """ | |
| if not sample: | |
| return torch.clamp(mu, -Z_CLAMP, Z_CLAMP) | |
| std = torch.exp(0.5 * logvar) | |
| std = torch.clamp(std, max=STD_MAX) | |
| eps = torch.randn_like(std) | |
| z = mu + eps * std | |
| z = torch.clamp(z, -Z_CLAMP, Z_CLAMP) | |
| return z | |
| def forward(self, x, label, sample=True): | |
| mu, logvar = self.encoder(x) | |
| z = self.reparametrize(mu, logvar, sample=sample) | |
| x_logits = self.decoder(z, label) | |
| return x_logits, mu, logvar | |
| def loss(self, x, x_logits, mu, logvar, beta=None): | |
| """ | |
| ELBO con la escala original: | |
| BCEWithLogits sumada por imagen + beta * KL | |
| Esta es la escala de la corrida previa: | |
| recon_loss ≈ 26000–28000 | |
| kl_loss ≈ 100–130 | |
| Usar con --no_amp. | |
| """ | |
| if beta is None: | |
| beta = self.beta | |
| B = x.size(0) | |
| x_f = x.float() | |
| logits_f = x_logits.float() | |
| mu_f = mu.float() | |
| logvar_f = logvar.float() | |
| recon_loss = F.binary_cross_entropy_with_logits( | |
| logits_f, | |
| x_f, | |
| reduction="sum", | |
| ) / B | |
| kl_loss = 0.5 * torch.sum( | |
| logvar_f.exp() + mu_f.pow(2) - 1.0 - logvar_f | |
| ) / B | |
| total = recon_loss + beta * kl_loss | |
| return total, recon_loss, kl_loss | |
| def generate(self, class_label, n=1, device="cpu", temperature=1.0): | |
| """ | |
| Genera imágenes sintéticas en [0, 1]. | |
| """ | |
| self.eval() | |
| label = torch.full( | |
| size=(n,), | |
| fill_value=int(class_label), | |
| dtype=torch.long, | |
| device=device, | |
| ) | |
| z = torch.randn(n, self.latent_dim, device=device) * temperature | |
| z = torch.clamp(z, -Z_CLAMP, Z_CLAMP) | |
| logits = self.decoder(z, label) | |
| return torch.sigmoid(logits) | |
| def reconstruct(self, x, label): | |
| """ | |
| Reconstrucción determinista usando mu. | |
| Devuelve imágenes en [0, 1]. | |
| """ | |
| self.eval() | |
| mu, logvar = self.encoder(x) | |
| z = self.reparametrize(mu, logvar, sample=False) | |
| logits = self.decoder(z, label) | |
| return torch.sigmoid(logits), mu | |
| def interpolate(self, x1, label1, x2, label2, steps=8, device="cpu"): | |
| """ | |
| Interpolación en espacio latente. | |
| Devuelve imágenes en [0, 1]. | |
| """ | |
| self.eval() | |
| mu1, _ = self.encoder(x1) | |
| mu2, _ = self.encoder(x2) | |
| alphas = torch.linspace(0, 1, steps, device=device) | |
| results = [] | |
| for a in alphas: | |
| z = (1.0 - a) * mu1 + a * mu2 | |
| z = torch.clamp(z, -Z_CLAMP, Z_CLAMP) | |
| logits = self.decoder(z, label1) | |
| results.append(torch.sigmoid(logits)) | |
| return torch.cat(results, dim=0) | |
| # ─── Test rápido ──────────────────────────────────────────────────────────── | |
| if __name__ == "__main__": | |
| device = "cuda" if torch.cuda.is_available() else "cpu" | |
| model = ConditionalVAE(latent_dim=128, num_classes=7, beta=1.0).to(device) | |
| x = torch.rand(4, 3, 128, 128, device=device) | |
| label = torch.randint(0, 7, (4,), device=device) | |
| x_logits, mu, logvar = model(x, label, sample=True) | |
| loss, recon, kl = model.loss(x, x_logits, mu, logvar) | |
| print(f"Input: {x.shape} min={x.min():.3f} max={x.max():.3f}") | |
| print(f"Logits: {x_logits.shape} min={x_logits.min():.3f} max={x_logits.max():.3f}") | |
| print(f"mu: min={mu.min():.3f} max={mu.max():.3f}") | |
| print(f"logvar: min={logvar.min():.3f} max={logvar.max():.3f}") | |
| print(f"Loss: total={loss:.4f} recon={recon:.4f} kl={kl:.4f}") | |
| assert torch.isfinite(loss), "NaN/Inf en loss" | |
| assert torch.isfinite(x_logits).all(), "NaN/Inf en logits" | |
| assert torch.isfinite(mu).all(), "NaN/Inf en mu" | |
| assert torch.isfinite(logvar).all(), "NaN/Inf en logvar" | |
| gen = model.generate(class_label=0, n=4, device=device) | |
| print(f"Generated: {gen.shape} min={gen.min():.3f} max={gen.max():.3f}") | |
| total_params = sum(p.numel() for p in model.parameters()) | |
| print(f"Parámetros: {total_params:,}") | |
| print("✅ Test pasado") | |