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