Calamar49 commited on
Commit
20c2bab
·
1 Parent(s): adc96d8

return functional model

Browse files
Files changed (2) hide show
  1. Dockerfile +1 -1
  2. app.py +16 -50
Dockerfile CHANGED
@@ -5,4 +5,4 @@ COPY ./requirements.txt /code/requirements.txt
5
  RUN pip install --no-cache-dir --upgrade -r /code/requirements.txt
6
  COPY . /code
7
  EXPOSE 7860
8
- CMD ["uvicorn", "app:app", "--host", "0.0.0.0", "--port", "7860"]
 
5
  RUN pip install --no-cache-dir --upgrade -r /code/requirements.txt
6
  COPY . /code
7
  EXPOSE 7860
8
+ CMD ["python3", "app.py"]
app.py CHANGED
@@ -1,4 +1,4 @@
1
- # app.py (versión final con Gradio y API para móvil)
2
 
3
  import torch
4
  import torchvision.transforms as transforms
@@ -6,9 +6,6 @@ from PIL import Image
6
  import json
7
  import timm
8
  import gradio as gr
9
- from fastapi import FastAPI, UploadFile, File
10
- from fastapi.responses import JSONResponse
11
- import io
12
 
13
  # --- 1. Importar la definición del modelo ---
14
  from VisionEnsembleModel import VisionEnsembleModel
@@ -19,18 +16,14 @@ MODEL_PATH = "model/best_vision_ensemble_model.pth"
19
  LABELS_PATH = "model/species_labels_map.json"
20
  NUM_CLASSES = 156
21
 
22
- try:
23
- with open(LABELS_PATH) as f:
24
- labels_map = json.load(f)
25
- print("Mapa de etiquetas cargado.")
26
- except Exception as e:
27
- labels_map = {}
28
- print(f"Error cargando el mapa de etiquetas: {e}")
29
 
30
  model = VisionEnsembleModel(num_classes=NUM_CLASSES)
31
  model.load_state_dict(torch.load(MODEL_PATH, map_location=device))
32
  model.to(device)
33
  model.eval()
 
34
  print("Modelo Ensamblado Híbrido (CNN+ViT) cargado y listo.")
35
 
36
  transforms_val = transforms.Compose([
@@ -39,60 +32,33 @@ transforms_val = transforms.Compose([
39
  transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
40
  ])
41
 
42
- # --- 3. Función de Predicción ---
43
- def make_prediction(image_pil):
44
- """Función interna que toma una imagen PIL y devuelve las probabilidades."""
45
- input_tensor = transforms_val(image_pil).unsqueeze(0).to(device)
46
  with torch.no_grad():
47
  output = model(input_tensor)
48
  probabilities = torch.nn.functional.softmax(output[0], dim=0)
49
- return probabilities
50
 
51
- # --- 4. Lógica para la Interfaz de Gradio (devuelve nombres) ---
52
- def predict_for_gradio(image_numpy):
53
- pil_image = Image.fromarray(image_numpy.astype('uint8'), 'RGB')
54
- probabilities = make_prediction(pil_image)
55
-
56
  top5_prob, top5_catid = torch.topk(probabilities, 5)
57
  confidences = {}
58
  for i in range(top5_prob.size(0)):
59
  species_id = top5_catid[i].item()
60
  prob = top5_prob[i].item()
61
- species_name = labels_map.get(str(species_id), f"ID Desconocido: {species_id}")
62
  confidences[species_name] = prob
63
  return confidences
64
 
65
- # --- 5. Crear la Interfaz de Gradio ---
66
- gradio_interface = gr.Interface(
67
- fn=predict_for_gradio,
68
  inputs=gr.Image(type="numpy", label="Sube una imagen de tu orquídea"),
69
  outputs=gr.Label(num_top_classes=5, label="Predicciones"),
70
  title="Clasificador de Orquídeas (Modelo Ensamblado)",
71
  description="Sube una foto de una orquídea y la IA (CNN+ViT) intentará identificar la especie.",
72
  )
73
 
74
- # --- 6. Crear la App FastAPI ---
75
- app = FastAPI()
76
-
77
- # --- 7. Endpoint de API para la App Móvil (devuelve IDs) ---
78
- @app.post("/predict_for_mobile")
79
- async def predict_for_mobile(file: UploadFile = File(...)):
80
- image_bytes = await file.read()
81
- try:
82
- pil_image = Image.open(io.BytesIO(image_bytes)).convert("RGB")
83
- except Exception:
84
- return JSONResponse(status_code=400, content={"error": "Archivo de imagen inválido."})
85
-
86
- probabilities = make_prediction(pil_image)
87
- top5_prob, top_catid = torch.topk(probabilities, 5)
88
-
89
- results = []
90
- for i in range(top5_prob.size(0)):
91
- results.append({
92
- "species_id": top_catid[i].item(),
93
- "confidence": top_prob[i].item()
94
- })
95
- return JSONResponse(content={"predictions": results})
96
-
97
- # --- 8. Montar la Interfaz de Gradio en la ruta raíz ---
98
- app = gr.mount_gradio_app(app, gradio_interface, path="/")
 
1
+ # app.py (versión final y simplificada con Gradio)
2
 
3
  import torch
4
  import torchvision.transforms as transforms
 
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
 
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
  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)