""" 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 @torch.no_grad() 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) @torch.no_grad() 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 @torch.no_grad() 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")