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

Final architecture: Gradio UI + FastAPI endpoint for mobile

Browse files
Files changed (2) hide show
  1. Dockerfile +2 -1
  2. app.py +50 -16
Dockerfile CHANGED
@@ -1,7 +1,8 @@
 
1
  FROM huggingface/transformers-pytorch-gpu
2
  WORKDIR /code
3
  COPY ./requirements.txt /code/requirements.txt
4
  RUN pip install --no-cache-dir --upgrade -r /code/requirements.txt
5
  COPY . /code
6
  EXPOSE 7860
7
- CMD ["python3", "app.py"]
 
1
+ # Contenido de: Dockerfile
2
  FROM huggingface/transformers-pytorch-gpu
3
  WORKDIR /code
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"]
app.py CHANGED
@@ -1,4 +1,4 @@
1
- # app.py (versión final y simplificada con Gradio)
2
 
3
  import torch
4
  import torchvision.transforms as transforms
@@ -6,6 +6,9 @@ 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
@@ -16,14 +19,18 @@ 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,33 +39,60 @@ 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)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # app.py (versión final con Gradio y API para móvil)
2
 
3
  import torch
4
  import torchvision.transforms as transforms
 
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
  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
  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="/")