Calamar49's picture
translate
68d588b
Raw
History Blame Contribute Delete
4.15 kB
# 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()