DeepLearning / app.py
Chencho98's picture
Upload folder using huggingface_hub
63e3bc6 verified
Raw
History Blame Contribute Delete
3.71 kB
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()