import torch import gradio as gr import numpy as np from PIL import Image from model import ConditionalVAE CHECKPOINT_PATH = "best_model_inference.pt" DEVICE = "cuda" if torch.cuda.is_available() else "cpu" DEFAULT_CLASS_NAMES = ["akiec", "bcc", "bkl", "df", "mel", "nv", "vasc"] CLASS_DESCRIPTIONS = { "akiec": "AKIEC — Queratosis actínica / enfermedad de Bowen", "bcc": "BCC — Carcinoma basocelular", "bkl": "BKL — Lesión benigna tipo queratosis", "df": "DF — Dermatofibroma", "mel": "MEL — Melanoma", "nv": "NV — Nevus melanocítico", "vasc": "VASC — Lesión vascular", } def load_cvae(): ckpt = torch.load(CHECKPOINT_PATH, map_location=DEVICE, weights_only=False) class_names = ckpt.get("class_names", DEFAULT_CLASS_NAMES) args = ckpt.get("args", {}) latent_dim = int(args.get("latent_dim", 128)) beta = float(args.get("beta", 1.0)) model = ConditionalVAE( latent_dim=latent_dim, num_classes=len(class_names), beta=beta, ).to(DEVICE) model.load_state_dict(ckpt["model"]) model.eval() return model, class_names, ckpt model, class_names, ckpt = load_cvae() def tensor_to_pil(img_tensor): """ img_tensor: Tensor (3, H, W) en [0, 1] """ img = img_tensor.detach().cpu().clamp(0, 1) img = img.permute(1, 2, 0).numpy() img = (img * 255).astype(np.uint8) return Image.fromarray(img) @torch.no_grad() def generate_images(class_name, temperature, n_images): class_idx = class_names.index(class_name) imgs = model.generate( class_label=class_idx, n=int(n_images), device=DEVICE, temperature=float(temperature), ) return [tensor_to_pil(imgs[i]) for i in range(imgs.size(0))] def class_info(class_name): return CLASS_DESCRIPTIONS.get(class_name, class_name) description = f""" # CVAE — Generación sintética de lesiones de piel Este demo genera imágenes sintéticas de lesiones dermatológicas usando un Conditional VAE entrenado sobre 7 clases. **Modelo:** Conditional Variational Autoencoder **Checkpoint epoch:** {ckpt.get("epoch", "N/A")} **Best val loss:** {ckpt.get("best_val", "N/A")} **Clases:** {", ".join(class_names)} ⚠️ Demo académico. No usar para diagnóstico médico. """ with gr.Blocks(title="Skin Lesion CVAE") as demo: gr.Markdown(description) with gr.Row(): with gr.Column(): class_name = gr.Dropdown( choices=class_names, value=class_names[0], label="Clase de lesión", ) class_description = gr.Textbox( value=class_info(class_names[0]), label="Descripción", interactive=False, ) temperature = gr.Slider( minimum=0.2, maximum=1.5, value=0.8, step=0.1, label="Temperatura", ) n_images = gr.Slider( minimum=1, maximum=8, value=4, step=1, label="Número de imágenes", ) btn = gr.Button("Generar imágenes") with gr.Column(): gallery = gr.Gallery( label="Imágenes sintéticas generadas", columns=4, height="auto", ) class_name.change( fn=class_info, inputs=class_name, outputs=class_description, ) btn.click( fn=generate_images, inputs=[class_name, temperature, n_images], outputs=gallery, ) if __name__ == "__main__": demo.launch()