Calamar49 commited on
Commit
f28b4c0
·
1 Parent(s): 4937f30

return best_model

Browse files
Files changed (2) hide show
  1. VisionEnsembleModel.py +0 -6
  2. app.py +15 -33
VisionEnsembleModel.py CHANGED
@@ -5,14 +5,8 @@ import torch.nn as nn
5
  import timm
6
 
7
  class VisionEnsembleModel(nn.Module):
8
- """
9
- La misma clase de modelo que definiste en Colab.
10
- """
11
  def __init__(self, num_classes, cnn_model_name='efficientnet_b2', vit_model_name='vit_small_patch16_224'):
12
  super().__init__()
13
- # Usamos pretrained=False aquí porque cargaremos nuestros propios pesos.
14
- # Timm cargará los pesos preentrenados si no encuentra un state_dict local,
15
- # pero es más limpio ser explícito. Al final, los sobrescribiremos.
16
  self.cnn = timm.create_model(cnn_model_name, pretrained=False, num_classes=num_classes)
17
  cnn_features = self.cnn.get_classifier().in_features
18
  self.cnn.reset_classifier(0)
 
5
  import timm
6
 
7
  class VisionEnsembleModel(nn.Module):
 
 
 
8
  def __init__(self, num_classes, cnn_model_name='efficientnet_b2', vit_model_name='vit_small_patch16_224'):
9
  super().__init__()
 
 
 
10
  self.cnn = timm.create_model(cnn_model_name, pretrained=False, num_classes=num_classes)
11
  cnn_features = self.cnn.get_classifier().in_features
12
  self.cnn.reset_classifier(0)
app.py CHANGED
@@ -1,16 +1,19 @@
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_vision_ensemble_model.pth" # Usaremos el modelo ligero que sabemos que funciona
 
 
 
14
  LABELS_PATH = "model/species_labels_map.json"
15
  NUM_CLASSES = 156
16
 
@@ -18,65 +21,44 @@ NUM_CLASSES = 156
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()
 
1
+ # app.py (Versión Corregida para Cargar el Modelo Ensamblado)
2
 
3
  import torch
4
  import torchvision.transforms as transforms
5
  from PIL import Image
6
  import json
 
7
  import gradio as gr
8
 
9
+ # --- ¡CAMBIO 1: Importar la clase de nuestro modelo! ---
10
+ from VisionEnsembleModel import VisionEnsembleModel
11
 
12
+ # --- 1. Carga del Modelo y Componentes ---
13
+ device = torch.device("cpu")
14
+
15
+ # --- ¡VERIFICACIÓN! Apuntar al archivo de modelo correcto ---
16
+ MODEL_PATH = "model/best_vision_ensemble_model.pth"
17
  LABELS_PATH = "model/species_labels_map.json"
18
  NUM_CLASSES = 156
19
 
 
21
  with open(LABELS_PATH) as f:
22
  labels_map = json.load(f)
23
 
24
+ # --- ¡CAMBIO 2: Instanciar nuestro modelo personalizado! ---
25
+ model = VisionEnsembleModel(num_classes=NUM_CLASSES)
26
  model.load_state_dict(torch.load(MODEL_PATH, map_location=device))
27
  model.to(device)
28
  model.eval()
29
 
30
+ print("Modelo VisionEnsembleModel cargado y listo.")
31
 
32
+ # --- 2. Definir la Función de Predicción para Gradio (SIN CAMBIOS) ---
33
  transforms_val = transforms.Compose([
34
  transforms.Resize((224, 224)),
35
  transforms.ToTensor(),
36
  transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
37
  ])
38
 
 
 
39
  def predict(image):
 
 
 
 
40
  pil_image = Image.fromarray(image.astype('uint8'), 'RGB')
 
 
41
  input_tensor = transforms_val(pil_image).unsqueeze(0).to(device)
 
 
42
  with torch.no_grad():
43
  output = model(input_tensor)
44
  probabilities = torch.nn.functional.softmax(output[0], dim=0)
45
 
 
46
  top5_prob, top5_catid = torch.topk(probabilities, 5)
 
47
  confidences = {}
48
  for i in range(top5_prob.size(0)):
49
  species_id = top5_catid[i].item()
50
  prob = top5_prob[i].item()
51
  species_name = labels_map.get(str(species_id), "Desconocido")
52
  confidences[species_name] = prob
 
53
  return confidences
54
 
55
+ # --- 3. Crear y Lanzar la Interfaz de Gradio (SIN CAMBIOS) ---
 
 
56
  iface = gr.Interface(
57
  fn=predict,
58
  inputs=gr.Image(type="numpy", label="Sube una imagen de tu orquídea"),
59
  outputs=gr.Label(num_top_classes=5, label="Predicciones"),
60
+ title="Clasificador de Orquídeas (Modelo Ensamblado)",
61
+ description="Sube una foto de una orquídea y la IA intentará identificar la especie.",
 
 
 
 
 
62
  )
63
 
 
64
  iface.launch()