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

pure model

Browse files
Files changed (1) hide show
  1. app.py +31 -60
app.py CHANGED
@@ -1,38 +1,27 @@
1
- # app.py (versión final con depuración)
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. 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([
38
  transforms.Resize((224, 224)),
@@ -40,48 +29,30 @@ 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"),
82
- outputs=gr.Label(num_top_classes=5, label="Predicciones"),
83
- title="Clasificador de Orquídeas (Modelo Ensamblado)",
84
- description="Sube una foto de una orquídea y la IA (CNN+ViT) intentará identificar la especie.",
85
- )
86
-
87
- iface.launch(server_name="0.0.0.0", server_port=7860)
 
1
+ # app.py (versión final - API Numérica)
2
 
3
  import torch
4
  import torchvision.transforms as transforms
5
  from PIL import Image
6
  import json
7
  import timm
8
+ from fastapi import FastAPI, UploadFile, File
9
+ from fastapi.responses import JSONResponse
10
+ import io
11
 
12
  # --- 1. Importar la definición del modelo ---
13
  from VisionEnsembleModel import VisionEnsembleModel
14
 
15
+ # --- 2. Carga del Modelo ---
16
  device = torch.device("cpu")
17
  MODEL_PATH = "model/best_vision_ensemble_model.pth"
 
18
  NUM_CLASSES = 156
19
 
 
 
 
 
 
 
 
 
 
 
 
 
20
  model = VisionEnsembleModel(num_classes=NUM_CLASSES)
21
  model.load_state_dict(torch.load(MODEL_PATH, map_location=device))
22
  model.to(device)
23
  model.eval()
24
+ print("Modelo Ensamblado Híbrido cargado.")
25
 
26
  transforms_val = transforms.Compose([
27
  transforms.Resize((224, 224)),
 
29
  transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
30
  ])
31
 
32
+ # --- 3. Crear la App FastAPI ---
33
+ app = FastAPI(title="API Numérica de Clasificación de Orquídeas")
34
+
35
+ # --- 4. Endpoint de Predicción Numérica ---
36
+ @app.post("/predict_numeric")
37
+ async def predict_numeric(file: UploadFile = File(...)):
38
+ image_bytes = await file.read()
39
  try:
40
+ image = Image.open(io.BytesIO(image_bytes)).convert("RGB")
41
+ except Exception:
42
+ return JSONResponse(status_code=400, content={"error": "Archivo de imagen inválido."})
43
+
44
+ input_tensor = transforms_val(image).unsqueeze(0).to(device)
45
+ with torch.no_grad():
46
+ output = model(input_tensor)
47
+ probabilities = torch.nn.functional.softmax(output[0], dim=0)
48
+
49
+ top5_prob, top5_catid = torch.topk(probabilities, 5)
50
+
51
+ results = []
52
+ for i in range(top5_prob.size(0)):
53
+ results.append({
54
+ "species_id": top5_catid[i].item(),
55
+ "confidence": top5_prob[i].item()
56
+ })
57
 
58
+ return JSONResponse(content={"predictions": results})