Calamar49 commited on
Commit
35ed367
·
1 Parent(s): 168ce67

return version functional

Browse files
Files changed (1) hide show
  1. app.py +34 -28
app.py CHANGED
@@ -1,27 +1,30 @@
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,30 +32,33 @@ transforms_val = transforms.Compose([
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})
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # app.py (versión final y simplificada 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. 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
+ 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([
30
  transforms.Resize((224, 224)),
 
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 (sin cambios) ---
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
+ confidences = {}
 
45
  for i in range(top5_prob.size(0)):
46
+ species_id = top5_catid[i].item()
47
+ prob = top5_prob[i].item()
48
+ species_name = labels_map.get(str(species_id), "Desconocido")
49
+ confidences[species_name] = prob
50
+ return confidences
51
+
52
+ # --- 4. Crear y Lanzar la Interfaz de Gradio ---
53
+ iface = gr.Interface(
54
+ fn=predict,
55
+ inputs=gr.Image(type="numpy", label="Sube una imagen de tu orquídea"),
56
+ outputs=gr.Label(num_top_classes=5, label="Predicciones"),
57
+ title="Clasificador de Orquídeas (Modelo Ensamblado)",
58
+ description="Sube una foto de una orquídea y la IA (CNN+ViT) intentará identificar la especie.",
59
+ )
60
+
61
+ # Lanzamos la aplicación.
62
+ # server_name="0.0.0.0" es crucial para que funcione dentro de Docker.
63
+ # server_port=7860 es el puerto estándar que Hugging Face expone.
64
+ iface.launch(server_name="0.0.0.0", server_port=7860)