Calamar49 commited on
Commit
3b989b7
·
1 Parent(s): 92bdda0

Migrate to a Gradio interface for user-friendly predictions

Browse files
Files changed (4) hide show
  1. Dockerfile +0 -22
  2. README.md +7 -12
  3. app.py +49 -56
  4. requirements.txt +2 -1
Dockerfile DELETED
@@ -1,22 +0,0 @@
1
- # Usa la imagen base oficial de Hugging Face para Spaces. Incluye PyTorch y CUDA.
2
- FROM huggingface/transformers-pytorch-gpu
3
-
4
- # Establece el directorio de trabajo dentro del contenedor
5
- WORKDIR /code
6
-
7
- # Copia el archivo de requerimientos primero para que Docker pueda cachear la instalación
8
- COPY ./requirements.txt /code/requirements.txt
9
-
10
- # Instala todas las dependencias de Python
11
- RUN pip install --no-cache-dir --upgrade -r /code/requirements.txt
12
-
13
- # Copia todos los demás archivos de tu proyecto al contenedor
14
- COPY . /code
15
-
16
- # Expone el puerto que usará la aplicación. 7860 es el estándar para Spaces.
17
- EXPOSE 7860
18
-
19
- # Define el comando que se ejecutará para iniciar tu API de FastAPI.
20
- # Le dice a 'uvicorn' que busque un objeto llamado 'app' en un archivo llamado 'app.py'
21
- # y que lo sirva en todas las interfaces de red ('0.0.0.0') en el puerto 7860.
22
- CMD ["uvicorn", "app:app", "--host", "0.0.0.0", "--port", "7860"]
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
README.md CHANGED
@@ -1,14 +1,9 @@
1
  ---
2
- title: Orchid Classifier Api
3
- emoji: 🌸 # Cambié el emoji a una flor, ¡más temático!
4
- colorFrom: purple
5
- colorTo: blue
6
- sdk: docker
7
- app_file: app.py
8
- app_port: 7860
9
- pinned: false
10
  license: mit
11
- short_description: api
12
- ---
13
-
14
- Check out the configuration reference at https://huggingface.co/docs/hub/spaces-config-reference
 
1
  ---
2
+ title: Orchid Classifier
3
+ emoji: 🌸
4
+ colorFrom: green
5
+ colorTo: purple
6
+ sdk: gradio
7
+ sdk_version: 4.31.0 # Usar una versión reciente de Gradio
 
 
8
  license: mit
9
+ ---
 
 
 
app.py CHANGED
@@ -1,89 +1,82 @@
1
- # Contenido de: app.py
2
 
3
  import torch
4
  import torchvision.transforms as transforms
5
  from PIL import Image
6
  import json
7
- from fastapi import FastAPI, UploadFile, File
8
- from fastapi.responses import JSONResponse
9
- import io
10
 
11
- # Importamos la definición de nuestro modelo desde el otro archivo
12
- from VisionEnsembleModel import VisionEnsembleModel
13
 
14
- # --- 1. Carga del Modelo y Componentes ---
15
-
16
- # Definimos el dispositivo (en los Spaces de Hugging Face, podemos usar CPU o GPU)
17
- device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
18
- print(f"Usando dispositivo: {device}")
19
-
20
- # Rutas a los archivos
21
- MODEL_PATH = "model/best_vision_ensemble_model.pth"
22
  LABELS_PATH = "model/species_labels_map.json"
23
- NUM_CLASSES = 156 # El número de clases con el que fue entrenado
24
 
25
- # Cargamos el mapa de etiquetas
26
  with open(LABELS_PATH) as f:
27
  labels_map = json.load(f)
28
- print("Mapa de etiquetas cargado.")
29
 
30
- # Instanciamos el modelo y cargamos los pesos
31
- model = VisionEnsembleModel(num_classes=NUM_CLASSES)
32
  model.load_state_dict(torch.load(MODEL_PATH, map_location=device))
33
  model.to(device)
34
- model.eval() # ¡Muy importante poner el modelo en modo de evaluación!
35
- print("Modelo cargado y en modo de evaluación.")
 
36
 
37
- # Definimos las transformaciones de la imagen (deben ser las mismas que en la validación)
38
  transforms_val = transforms.Compose([
39
  transforms.Resize((224, 224)),
40
  transforms.ToTensor(),
41
  transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
42
  ])
43
 
44
- # --- 2. Creación de la Aplicación FastAPI ---
45
-
46
- app = FastAPI(title="API de Clasificación de Orquídeas")
47
-
48
- @app.get("/")
49
- def read_root():
50
- return {"message": "Bienvenido a la API de Orquídeas. Envía una imagen al endpoint /predict"}
51
 
52
- @app.post("/predict")
53
- async def predict(file: UploadFile = File(...)):
54
  """
55
- Endpoint que recibe una imagen, la procesa y devuelve la predicción.
56
  """
57
- # Leer el contenido de la imagen en memoria
58
- image_bytes = await file.read()
59
- try:
60
- image = Image.open(io.BytesIO(image_bytes)).convert("RGB")
61
- except Exception as e:
62
- return JSONResponse(status_code=400, content={"error": f"Archivo inválido: {e}"})
63
-
64
  # Preprocesar la imagen
65
- input_tensor = transforms_val(image).unsqueeze(0).to(device)
66
 
67
  # Realizar la predicción
68
  with torch.no_grad():
69
  output = model(input_tensor)
70
  probabilities = torch.nn.functional.softmax(output[0], dim=0)
71
 
72
- # Obtener la predicción principal
73
- top_prob, top_catid = torch.topk(probabilities, 1)
74
- predicted_id = top_catid[0].item()
75
- confidence = top_prob[0].item()
76
 
77
- # Traducir el ID a un nombre de especie
78
- predicted_species = labels_map.get(str(predicted_id), "Especie Desconocida")
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
79
 
80
- # Devolver el resultado en formato JSON
81
- return JSONResponse(
82
- status_code=200,
83
- content={
84
- "filename": file.filename,
85
- "predicted_species": predicted_species,
86
- "confidence": f"{confidence:.4f}",
87
- "species_id": predicted_id
88
- }
89
- )
 
1
+ # app.py (versión con Gradio)
2
 
3
  import torch
4
  import torchvision.transforms as transforms
5
  from PIL import Image
6
  import json
7
+ import timm
8
+ import gradio as gr
 
9
 
10
+ # --- 1. Carga del Modelo y Componentes (igual que antes) ---
 
11
 
12
+ device = torch.device("cpu") # Es más seguro usar CPU en el plan gratuito
13
+ MODEL_PATH = "model/best_mobilenet_model.pth" # Usaremos el modelo ligero que sabemos que funciona
 
 
 
 
 
 
14
  LABELS_PATH = "model/species_labels_map.json"
15
+ NUM_CLASSES = 156
16
 
17
+ # Cargar mapa de etiquetas
18
  with open(LABELS_PATH) as f:
19
  labels_map = json.load(f)
 
20
 
21
+ # Definir y cargar el modelo ligero
22
+ model = timm.create_model('mobilenetv2_100', pretrained=False, num_classes=NUM_CLASSES)
23
  model.load_state_dict(torch.load(MODEL_PATH, map_location=device))
24
  model.to(device)
25
+ model.eval()
26
+
27
+ print("Modelo MobileNetV2 cargado y listo.")
28
 
29
+ # Definir las transformaciones de la imagen
30
  transforms_val = transforms.Compose([
31
  transforms.Resize((224, 224)),
32
  transforms.ToTensor(),
33
  transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
34
  ])
35
 
36
+ # --- 2. Definir la Función de Predicción para Gradio ---
 
 
 
 
 
 
37
 
38
+ def predict(image):
 
39
  """
40
+ Esta función toma una imagen (de Gradio) y devuelve un diccionario de predicciones.
41
  """
42
+ # La imagen de Gradio viene como un array de Numpy, la convertimos a PIL Image
43
+ pil_image = Image.fromarray(image.astype('uint8'), 'RGB')
44
+
 
 
 
 
45
  # Preprocesar la imagen
46
+ input_tensor = transforms_val(pil_image).unsqueeze(0).to(device)
47
 
48
  # Realizar la predicción
49
  with torch.no_grad():
50
  output = model(input_tensor)
51
  probabilities = torch.nn.functional.softmax(output[0], dim=0)
52
 
53
+ # Crear un diccionario de confianza para las 5 mejores predicciones
54
+ top5_prob, top5_catid = torch.topk(probabilities, 5)
 
 
55
 
56
+ confidences = {}
57
+ for i in range(top5_prob.size(0)):
58
+ species_id = top5_catid[i].item()
59
+ prob = top5_prob[i].item()
60
+ species_name = labels_map.get(str(species_id), "Desconocido")
61
+ confidences[species_name] = prob
62
+
63
+ return confidences
64
+
65
+ # --- 3. Crear y Lanzar la Interfaz de Gradio ---
66
+
67
+ # Creamos la interfaz
68
+ iface = gr.Interface(
69
+ fn=predict,
70
+ inputs=gr.Image(type="numpy", label="Sube una imagen de tu orquídea"),
71
+ outputs=gr.Label(num_top_classes=5, label="Predicciones"),
72
+ title="Clasificador de Orquídeas",
73
+ description="Sube una foto de una orquídea y la IA intentará identificar la especie. Este modelo usa un MobileNetV2.",
74
+ examples=[
75
+ # Puedes añadir rutas a imágenes de ejemplo si las subes a tu repositorio
76
+ # ["ejemplo1.jpg"],
77
+ # ["ejemplo2.jpg"]
78
+ ]
79
+ )
80
 
81
+ # Lanzamos la aplicación
82
+ iface.launch()
 
 
 
 
 
 
 
 
requirements.txt CHANGED
@@ -7,4 +7,5 @@ torch
7
  torchvision
8
  timm
9
  Pillow
10
- scikit-learn
 
 
7
  torchvision
8
  timm
9
  Pillow
10
+ scikit-learn
11
+ gradio