Calamar49 commited on
Commit
3c8c844
·
1 Parent(s): 3ea120a

fix id name

Browse files
Files changed (1) hide show
  1. app.py +45 -30
app.py CHANGED
@@ -1,4 +1,4 @@
1
- # app.py (versión final con nombres de especies en la salida)
2
 
3
  import torch
4
  import torchvision.transforms as transforms
@@ -10,20 +10,28 @@ import gradio as gr
10
  # --- 1. Importar la definición del modelo ---
11
  from VisionEnsembleModel import VisionEnsembleModel
12
 
13
- # --- 2. Carga del Modelo y Componentes (Sin cambios) ---
14
  device = torch.device("cpu")
15
  MODEL_PATH = "model/best_vision_ensemble_model.pth"
16
  LABELS_PATH = "model/species_labels_map.json"
17
  NUM_CLASSES = 156
18
 
19
- with open(LABELS_PATH) as f:
20
- labels_map = json.load(f)
 
 
 
 
 
 
 
 
 
21
 
22
  model = VisionEnsembleModel(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 Ensamblado Híbrido (CNN+ViT) cargado y listo.")
28
 
29
  transforms_val = transforms.Compose([
@@ -32,35 +40,42 @@ transforms_val = transforms.Compose([
32
  transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
33
  ])
34
 
35
- # --- 3. Función de Predicción (CON LA CORRECCIÓN) ---
36
  def predict(image):
37
- pil_image = Image.fromarray(image.astype('uint8'), 'RGB')
38
- input_tensor = transforms_val(pil_image).unsqueeze(0).to(device)
39
- with torch.no_grad():
40
- output = model(input_tensor)
41
- probabilities = torch.nn.functional.softmax(output[0], dim=0)
42
-
43
- top5_prob, top5_catid = torch.topk(probabilities, 5)
44
-
45
- # --- ¡ESTA ES LA CORRECCIÓN CLAVE! ---
46
- confidences = {}
47
- for i in range(top5_prob.size(0)):
48
- # Obtenemos el ID numérico predicho
49
- species_id = top5_catid[i].item()
50
- # Obtenemos la probabilidad
51
- prob = top5_prob[i].item()
52
 
53
- # Usamos el mapa de etiquetas para "traducir" el ID a un nombre.
54
- # Lo convertimos a string (str(species_id)) para que coincida con las claves del JSON.
55
- species_name = labels_map.get(str(species_id), f"Desconocido (ID: {species_id})")
56
 
57
- # Añadimos al diccionario el NOMBRE como clave y la probabilidad como valor.
58
- confidences[species_name] = prob
59
- # ------------------------------------
60
-
61
- return confidences
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
62
 
63
- # --- 4. Crear y Lanzar la Interfaz de Gradio (Sin cambios) ---
64
  iface = gr.Interface(
65
  fn=predict,
66
  inputs=gr.Image(type="numpy", label="Sube una imagen de tu orquídea"),
 
1
+ # app.py (versión final con depuración)
2
 
3
  import torch
4
  import torchvision.transforms as transforms
 
10
  # --- 1. Importar la definición del modelo ---
11
  from VisionEnsembleModel import VisionEnsembleModel
12
 
13
+ # --- 2. Carga del Modelo y Componentes ---
14
  device = torch.device("cpu")
15
  MODEL_PATH = "model/best_vision_ensemble_model.pth"
16
  LABELS_PATH = "model/species_labels_map.json"
17
  NUM_CLASSES = 156
18
 
19
+ try:
20
+ with open(LABELS_PATH) as f:
21
+ # Cargamos el mapa de etiquetas. Las claves JSON siempre son strings.
22
+ labels_map = json.load(f)
23
+ print("Mapa de etiquetas cargado con éxito.")
24
+ # Imprimimos una muestra para verificar
25
+ print("Ejemplo del mapa de etiquetas:", dict(list(labels_map.items())[:3]))
26
+ except Exception as e:
27
+ print(f"ERROR AL CARGAR EL MAPA DE ETIQUETAS: {e}")
28
+ labels_map = {}
29
+
30
 
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()
 
35
  print("Modelo Ensamblado Híbrido (CNN+ViT) cargado y listo.")
36
 
37
  transforms_val = transforms.Compose([
 
40
  transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
41
  ])
42
 
43
+ # --- 3. Función de Predicción (CON DEPURACIÓN) ---
44
  def predict(image):
45
+ print("\n--- Nueva Predicción Iniciada ---")
46
+ try:
47
+ pil_image = Image.fromarray(image.astype('uint8'), 'RGB')
48
+ input_tensor = transforms_val(pil_image).unsqueeze(0).to(device)
49
+ with torch.no_grad():
50
+ output = model(input_tensor)
51
+ probabilities = torch.nn.functional.softmax(output[0], dim=0)
 
 
 
 
 
 
 
 
52
 
53
+ top5_prob, top5_catid = torch.topk(probabilities, 5)
 
 
54
 
55
+ confidences = {}
56
+ for i in range(top5_prob.size(0)):
57
+ species_id = top5_catid[i].item()
58
+ prob = top5_prob[i].item()
59
+
60
+ # Líneas de depuración que veremos en los logs
61
+ print(f"Predicción {i+1}: ID numérico = {species_id} (Tipo: {type(species_id)})")
62
+
63
+ # Buscamos la clave como string
64
+ species_name = labels_map.get(str(species_id), f"ID Desconocido: {species_id}")
65
+
66
+ print(f"Nombre traducido: {species_name}")
67
+
68
+ confidences[species_name] = prob
69
+
70
+ print("--- Predicción completada con éxito ---")
71
+ return confidences
72
+ except Exception as e:
73
+ print(f"!!! ERROR DURANTE LA PREDICCIÓN: {e}")
74
+ # Devolvemos el error a la interfaz de Gradio para verlo
75
+ return {"Error": str(e)}
76
+
77
 
78
+ # --- 4. Crear y Lanzar la Interfaz de Gradio ---
79
  iface = gr.Interface(
80
  fn=predict,
81
  inputs=gr.Image(type="numpy", label="Sube una imagen de tu orquídea"),