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