File size: 3,714 Bytes
63e3bc6
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
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()