Spaces:
Sleeping
Sleeping
| 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) | |
| 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() | |