File size: 4,147 Bytes
7d0f1a1
5200fec
 
11ad57f
5200fec
 
 
 
11ad57f
5200fec
7d0f1a1
11ad57f
 
 
7d0f1a1
11ad57f
 
f32accf
11ad57f
 
 
f32accf
11ad57f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
5200fec
 
f32accf
5200fec
11ad57f
5200fec
 
11ad57f
7d0f1a1
11ad57f
 
 
f32accf
 
5200fec
 
 
 
f32accf
 
5200fec
 
 
 
 
 
 
68d588b
 
11ad57f
68d588b
11ad57f
 
 
5200fec
 
 
 
f32accf
11ad57f
68d588b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
11ad57f
 
f32accf
7d0f1a1
11ad57f
 
 
 
 
7d0f1a1
11ad57f
 
 
7d0f1a1
 
 
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
# app.py (Versión final, autocontenida y robusta para Gradio SDK)

import torch
import torch.nn as nn
import torchvision.transforms as transforms
from PIL import Image
import json
import timm
import gradio as gr

# --- 1. Definición del Modelo (directamente aquí para evitar errores) ---
class VisionEnsembleModel(nn.Module):
    def __init__(self, num_classes, cnn_model_name='efficientnet_b2', vit_model_name='vit_small_patch16_224'):
        super().__init__()
        # Se crean con pretrained=False porque cargaremos nuestros propios pesos
        self.cnn = timm.create_model(cnn_model_name, pretrained=False, num_classes=num_classes)
        cnn_features = self.cnn.get_classifier().in_features
        self.cnn.reset_classifier(0)

        self.vit = timm.create_model(vit_model_name, pretrained=False, num_classes=num_classes)
        vit_features = self.vit.head.in_features
        self.vit.head = nn.Identity()

        self.classifier = nn.Sequential(
            nn.BatchNorm1d(cnn_features + vit_features),
            nn.Linear(cnn_features + vit_features, 512),
            nn.ReLU(),
            nn.Dropout(0.5),
            nn.Linear(512, num_classes)
        )

    def forward(self, image):
        cnn_feat = self.cnn(image)
        vit_feat = self.vit(image)
        combined = torch.cat([cnn_feat, vit_feat], dim=1)
        output = self.classifier(combined)
        return output

# --- 2. Carga del Modelo y Componentes ---
device = torch.device("cpu")
MODEL_PATH = "model/best_vision_ensemble_model.pth"
LABELS_PATH = "model/species_labels_map.json"
NUM_CLASSES = 156

try:
    with open(LABELS_PATH, encoding="utf-8") as f:
        labels_map = json.load(f)
    print("Mapa de etiquetas cargado con éxito.")
except Exception as e:
    print(f"ERROR AL CARGAR MAPA DE ETIQUETAS: {e}")
    labels_map = {}

model = VisionEnsembleModel(num_classes=NUM_CLASSES)
model.load_state_dict(torch.load(MODEL_PATH, map_location=device))
model.to(device)
model.eval()
print("Modelo Híbrido cargado y listo.")

transforms_val = transforms.Compose([
    transforms.Resize((224, 224)),
    transforms.ToTensor(),
    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
])


# --- 3. Función de Predicción (CON LA TRADUCCIÓN A NOMBRES) ---
def predict(image):
    # La primera parte de la función no cambia
    if image is None:
        return None
    pil_image = Image.fromarray(image.astype('uint8'), 'RGB')
    input_tensor = transforms_val(pil_image).unsqueeze(0).to(device)
    with torch.no_grad():
        output = model(input_tensor)
        probabilities = torch.nn.functional.softmax(output[0], dim=0)

    top5_prob, top5_catid = torch.topk(probabilities, 5)
    
    # --- ¡AQUÍ ESTÁ LA ÚNICA CORRECCIÓN QUE NECESITAS! ---
    confidences = {}
    for i in range(top5_prob.size(0)):
        # Obtenemos el ID numérico predicho
        species_id = top5_catid[i].item()
        # Obtenemos la probabilidad
        prob = top5_prob[i].item()
        
        # Usamos el mapa de etiquetas para "traducir" el ID a un nombre.
        # Lo convertimos a string (str(species_id)) para que coincida con las claves del JSON.
        species_name = labels_map.get(str(species_id), f"Desconocido (ID: {species_id})")
        
        # Añadimos al diccionario el NOMBRE como clave y la probabilidad como valor.
        confidences[species_name] = prob
    # --------------------------------------------------------
        
    return confidences

# --- 4. Crear la Interfaz de Gradio ---
# La plataforma de Hugging Face encontrará esta variable 'iface' y la lanzará automáticamente.
iface = gr.Interface(
    fn=predict,
    inputs=gr.Image(type="numpy", label="Sube una imagen de tu orquídea"),
    outputs=gr.Label(num_top_classes=5, label="Predicciones"),
    title="Clasificador de Orquídeas",
    description="Sube una foto de una orquídea y la IA (un ensamblado de CNN y Vision Transformer) intentará identificar la especie.",
    allow_flagging="never"
)

# --- 5. Lanzar la demo (opcional pero recomendado para pruebas locales) ---
if __name__ == "__main__":
    iface.launch()