Calamar49 commited on
Commit
069e16c
·
1 Parent(s): 20c2bab
Files changed (2) hide show
  1. Dockerfile +1 -1
  2. app.py +36 -25
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 ["python3", "app.py"]
 
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
@@ -23,8 +26,7 @@ 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,33 +34,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 (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 para web y FastAPI 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
 
26
  model.load_state_dict(torch.load(MODEL_PATH, map_location=device))
27
  model.to(device)
28
  model.eval()
29
+ print("Modelo Ensamblado Híbrido cargado.")
 
30
 
31
  transforms_val = transforms.Compose([
32
  transforms.Resize((224, 224)),
 
34
  transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
35
  ])
36
 
37
+ # --- 3. Función de Predicción Interna ---
38
+ def make_prediction(image_pil):
39
+ input_tensor = transforms_val(image_pil).unsqueeze(0).to(device)
 
40
  with torch.no_grad():
41
  output = model(input_tensor)
42
  probabilities = torch.nn.functional.softmax(output[0], dim=0)
43
+ return probabilities
44
 
45
+ # --- 4. Función para la Interfaz de Gradio (devuelve nombres) ---
46
+ def predict_for_gradio(image_numpy):
47
+ pil_image = Image.fromarray(image_numpy.astype('uint8'), 'RGB')
48
+ probabilities = make_prediction(pil_image)
49
  top5_prob, top5_catid = torch.topk(probabilities, 5)
50
+ confidences = {labels_map.get(str(cat_id.item()), "Desconocido"): prob.item() for prob, cat_id in zip(top5_prob, top5_catid)}
 
 
 
 
 
51
  return confidences
52
 
53
+ # --- 5. Crear la Interfaz de Gradio ---
54
+ gradio_interface = gr.Interface(fn=predict_for_gradio, inputs=gr.Image(type="numpy"), outputs=gr.Label(num_top_classes=5), title="Clasificador de Orquídeas")
55
+
56
+ # --- 6. Crear la App FastAPI ---
57
+ app = FastAPI()
58
+
59
+ # --- 7. Endpoint de API para la App Móvil (devuelve IDs) ---
60
+ @app.post("/predict_for_mobile")
61
+ async def predict_for_mobile(file: UploadFile = File(...)):
62
+ image_bytes = await file.read()
63
+ try:
64
+ pil_image = Image.open(io.BytesIO(image_bytes)).convert("RGB")
65
+ except Exception:
66
+ return JSONResponse(status_code=400, content={"error": "Archivo de imagen inválido."})
67
+
68
+ probabilities = make_prediction(pil_image)
69
+ top5_prob, top_catid = torch.topk(probabilities, 5)
70
+
71
+ results = [{"species_id": cat_id.item(), "confidence": prob.item()} for prob, cat_id in zip(top5_prob, top_catid)]
72
+ return JSONResponse(content={"predictions": results})
73
 
74
+ # --- 8. Montar Gradio en la App FastAPI ---
75
+ app = gr.mount_gradio_app(app, gradio_interface, path="/")