Calamar49 commited on
Commit
b6f8b84
·
1 Parent(s): 7578b2b

best_model

Browse files
Files changed (1) hide show
  1. app.py +15 -19
app.py CHANGED
@@ -1,4 +1,4 @@
1
- # app.py (versión final combinando Gradio y FastAPI)
2
 
3
  import torch
4
  import torchvision.transforms as transforms
@@ -8,63 +8,59 @@ import timm
8
  import gradio as gr
9
  from fastapi import FastAPI
10
 
11
- # --- 1. Carga del Modelo y Componentes (Sin cambios) ---
 
12
 
13
- device = torch.device("cpu")
14
- MODEL_PATH = "model/best_mobilenet_model.pth"
 
15
  LABELS_PATH = "model/species_labels_map.json"
16
  NUM_CLASSES = 156
17
 
18
  with open(LABELS_PATH) as f:
19
  labels_map = json.load(f)
20
 
21
- model = timm.create_model('mobilenetv2_100', pretrained=False, num_classes=NUM_CLASSES)
 
22
  model.load_state_dict(torch.load(MODEL_PATH, map_location=device))
23
  model.to(device)
24
  model.eval()
25
 
26
- print("Modelo MobileNetV2 cargado y listo.")
27
 
 
28
  transforms_val = transforms.Compose([
29
  transforms.Resize((224, 224)),
30
  transforms.ToTensor(),
31
  transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
32
  ])
33
 
34
- # --- 2. Función de Predicción (Sin cambios) ---
35
-
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
-
40
  with torch.no_grad():
41
  output = model(input_tensor)
42
  probabilities = torch.nn.functional.softmax(output[0], dim=0)
43
 
44
  top5_prob, top5_catid = torch.topk(probabilities, 5)
45
-
46
  confidences = {}
47
  for i in range(top5_prob.size(0)):
48
  species_id = top5_catid[i].item()
49
  prob = top5_prob[i].item()
50
  species_name = labels_map.get(str(species_id), "Desconocido")
51
  confidences[species_name] = prob
52
-
53
  return confidences
54
 
55
- # --- 3. Crear la Interfaz de Gradio (Sin 'launch()') ---
56
-
57
  iface = gr.Interface(
58
  fn=predict,
59
  inputs=gr.Image(type="numpy", label="Sube una imagen de tu orquídea"),
60
  outputs=gr.Label(num_top_classes=5, label="Predicciones"),
61
- title="Clasificador de Orquídeas",
62
- description="Sube una foto de una orquídea y la IA intentará identificar la especie.",
63
  )
64
 
65
- # --- 4. Crear la App FastAPI y Montar Gradio en ella ---
66
-
67
  app = FastAPI()
68
-
69
- # Montamos la interfaz de Gradio en la ruta raíz ("/") de nuestra aplicación FastAPI
70
  app = gr.mount_gradio_app(app, iface, path="/")
 
1
+ # app.py (versión para el modelo híbrido con Gradio y FastAPI)
2
 
3
  import torch
4
  import torchvision.transforms as transforms
 
8
  import gradio as gr
9
  from fastapi import FastAPI
10
 
11
+ # --- 1. Importar la definición del modelo ---
12
+ from VisionEnsembleModel import VisionEnsembleModel # <-- ¡Importante!
13
 
14
+ # --- 2. Carga del Modelo y Componentes ---
15
+ device = torch.device("cpu") # Usar CPU es más seguro en el plan gratuito
16
+ MODEL_PATH = "model/best_vision_ensemble_model.pth" # <-- RUTA AL MODELO HÍBRIDO
17
  LABELS_PATH = "model/species_labels_map.json"
18
  NUM_CLASSES = 156
19
 
20
  with open(LABELS_PATH) as f:
21
  labels_map = json.load(f)
22
 
23
+ # --- Instanciamos y cargamos el modelo ensamblado ---
24
+ model = VisionEnsembleModel(num_classes=NUM_CLASSES) # <-- Usamos nuestra clase personalizada
25
  model.load_state_dict(torch.load(MODEL_PATH, map_location=device))
26
  model.to(device)
27
  model.eval()
28
 
29
+ print("Modelo Ensamblado Híbrido (CNN+ViT) cargado y listo.")
30
 
31
+ # Definir las transformaciones de la imagen (sin cambios)
32
  transforms_val = transforms.Compose([
33
  transforms.Resize((224, 224)),
34
  transforms.ToTensor(),
35
  transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
36
  ])
37
 
38
+ # --- 3. Función de Predicción (sin cambios en la lógica) ---
 
39
  def predict(image):
40
  pil_image = Image.fromarray(image.astype('uint8'), 'RGB')
41
  input_tensor = transforms_val(pil_image).unsqueeze(0).to(device)
 
42
  with torch.no_grad():
43
  output = model(input_tensor)
44
  probabilities = torch.nn.functional.softmax(output[0], dim=0)
45
 
46
  top5_prob, top5_catid = torch.topk(probabilities, 5)
 
47
  confidences = {}
48
  for i in range(top5_prob.size(0)):
49
  species_id = top5_catid[i].item()
50
  prob = top5_prob[i].item()
51
  species_name = labels_map.get(str(species_id), "Desconocido")
52
  confidences[species_name] = prob
 
53
  return confidences
54
 
55
+ # --- 4. Crear la Interfaz de Gradio (sin cambios) ---
 
56
  iface = gr.Interface(
57
  fn=predict,
58
  inputs=gr.Image(type="numpy", label="Sube una imagen de tu orquídea"),
59
  outputs=gr.Label(num_top_classes=5, label="Predicciones"),
60
+ title="Clasificador de Orquídeas (Modelo Ensamblado)",
61
+ description="Sube una foto de una orquídea y la IA (CNN+ViT) intentará identificar la especie.",
62
  )
63
 
64
+ # --- 5. Crear la App FastAPI y Montar Gradio (sin cambios) ---
 
65
  app = FastAPI()
 
 
66
  app = gr.mount_gradio_app(app, iface, path="/")