Allisona commited on
Commit
30f9e27
verified
1 Parent(s): 3e3b227

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +25 -23
app.py CHANGED
@@ -4,33 +4,38 @@ import joblib
4
  import pandas as pd
5
  import numpy as np
6
  import os
7
- from sklearn.preprocessing import StandardScaler # Asegura que joblib pueda cargar el objeto
 
 
 
 
 
 
8
 
9
  # --- 1. Cargar objetos serializados ---
10
  try:
11
- # Carga segura de artefactos en el entorno de Hugging Face
12
  model = joblib.load(os.path.join(os.path.dirname(__file__), 'model.pkl'))
13
  preprocessor = joblib.load(os.path.join(os.path.dirname(__file__), 'preprocessor.pkl'))
14
  model_columns = joblib.load(os.path.join(os.path.dirname(__file__), 'model_columns.pkl'))
15
  except Exception as e:
16
- # Manejo de error si los archivos no se encuentran o est谩n corruptos
17
  print(f"Error al cargar artefactos: {e}")
18
  model = None
19
  preprocessor = None
20
  model_columns = []
21
 
22
  # --- 2. Funci贸n de Predicci贸n (N煤cleo de la API) ---
23
- # Los argumentos de entrada coinciden con las columnas de X antes del preprocesamiento
24
  def predict_ko_tko(F1_KD, F2_KD, F1_STR, F2_STR, F1_TD, F2_TD, F1_SUB, F2_SUB, Round,
25
  F1_acc, F2_acc, KD_diff, STR_diff, TD_diff, SUB_diff, Location,
26
- # Se incluyen las columnas OHE ya existentes en el CSV como inputs discretos
27
  wc_B, wc_C, wc_F, wc_Fl, wc_H, wc_LH, wc_L, wc_M, wc_O, wc_SH, wc_W,
28
  wc_WB, wc_WF, wc_WFl, wc_WS):
29
 
30
  if not model or not preprocessor:
31
- return "Error", "Modelo o preprocesador no cargado."
 
32
 
33
- # 1. Crear el DataFrame de entrada con las 31 columnas originales de X
34
  input_data = pd.DataFrame({
35
  'Fighter_1_KD': [F1_KD], 'Fighter_2_KD': [F2_KD], 'Fighter_1_STR': [F1_STR], 'Fighter_2_STR': [F2_STR],
36
  'Fighter_1_TD': [F1_TD], 'Fighter_2_TD': [F2_TD], 'Fighter_1_SUB': [F1_SUB], 'Fighter_2_SUB': [F2_SUB],
@@ -46,11 +51,10 @@ def predict_ko_tko(F1_KD, F2_KD, F1_STR, F2_STR, F1_TD, F2_TD, F1_SUB, F2_SUB, R
46
  "weight_class_Women's Flyweight": [wc_WFl], "weight_class_Women's Strawweight": [wc_WS]
47
  })
48
 
49
- # 2. Preprocesamiento: Utilizar el ColumnTransformer ajustado.
50
- # Esto escala las estad铆sticas y aplica OHE a 'Location'.
51
  X_processed = preprocessor.transform(input_data)
52
 
53
- # 3. Convertir a DataFrame y asegurar el orden de las columnas (CRUCIAL)
54
  X_final = pd.DataFrame(X_processed, columns=model_columns)
55
 
56
  # 4. Predicci贸n
@@ -63,7 +67,6 @@ def predict_ko_tko(F1_KD, F2_KD, F1_STR, F2_STR, F1_TD, F2_TD, F1_SUB, F2_SUB, R
63
  return prob_str, result_str
64
 
65
  # --- 3. Creaci贸n de la Interfaz Gradio ---
66
- # Definici贸n de inputs (Simplificado con valores por defecto)
67
  inputs = [
68
  gr.Slider(0, 5, value=1, step=1, label="KD P1"), gr.Slider(0, 5, value=0, step=1, label="KD P2"),
69
  gr.Slider(0, 300, value=70, label="STR P1"), gr.Slider(0, 300, value=50, label="STR P2"),
@@ -73,16 +76,15 @@ inputs = [
73
  gr.Slider(0, 1, value=0.4, label="Precisi贸n STR P1 (F1_acc)"), gr.Slider(0, 1, value=0.3, label="Precisi贸n STR P2 (F2_acc)"),
74
  gr.Slider(-5, 5, value=1, label="Diferencia de KD"), gr.Slider(-300, 300, value=20, label="Diferencia de STR"),
75
  gr.Slider(-20, 20, value=3, label="Diferencia de TD"), gr.Slider(-5, 5, value=0, label="Diferencia de SUB"),
76
- gr.Dropdown(df['Location'].unique().tolist(), value='Las Vegas, NV', label="Ubicaci贸n"),
77
- # Se a帽ade la categor铆a de peso como un switch binario (el usuario selecciona 1 y el resto 0)
78
- gr.Checkbox(value=True, label="Lightweight (wc_L)"), gr.Checkbox(value=False, label="Bantamweight (wc_B)"),
79
- gr.Checkbox(value=False, label="Catch Weight (wc_C)"), gr.Checkbox(value=False, label="Featherweight (wc_F)"),
80
- gr.Checkbox(value=False, label="Flyweight (wc_Fl)"), gr.Checkbox(value=False, label="Heavyweight (wc_H)"),
81
- gr.Checkbox(value=False, label="Light Heavyweight (wc_LH)"), gr.Checkbox(value=False, label="Middleweight (wc_M)"),
82
- gr.Checkbox(value=False, label="Open Weight (wc_O)"), gr.Checkbox(value=False, label="Super Heavyweight (wc_SH)"),
83
- gr.Checkbox(value=False, label="Welterweight (wc_W)"), gr.Checkbox(value=False, label="Women's Bantamweight (wc_WB)"),
84
- gr.Checkbox(value=False, label="Women's Featherweight (wc_WF)"), gr.Checkbox(value=False, label="Women's Flyweight (wc_WFl)"),
85
- gr.Checkbox(value=False, label="Women's Strawweight (wc_WS)")
86
  ]
87
 
88
  outputs = [gr.Textbox(label="Probabilidad de KO/TKO (%)"), gr.Textbox(label="Resultado M谩s Probable")]
@@ -92,5 +94,5 @@ gr.Interface(
92
  inputs=inputs,
93
  outputs=outputs,
94
  title="馃 Predictor de KO/TKO en Combates UFC (Despliegue ML)",
95
- description="Modelo Random Forest para predecir la finalizaci贸n de un combate. El modelo usa estad铆sticas de los peleadores y el contexto del evento."
96
- ).launch(server_name="0.0.0.0", server_port=7860)
 
4
  import pandas as pd
5
  import numpy as np
6
  import os
7
+ from sklearn.preprocessing import StandardScaler # Importado para evitar errores de unpickling
8
+
9
+ # --- CORRECCI脫N: Listas hardcodeadas para la interfaz (resuelve NameError) ---
10
+ UFC_LOCATIONS = [
11
+ 'Las Vegas, NV', 'Rio de Janeiro, Brazil', 'Abu Dhabi, UAE',
12
+ 'London, England', 'New York, NY', 'Outro' # Incluimos 'Outro' por si hay locations no vistas
13
+ ]
14
 
15
  # --- 1. Cargar objetos serializados ---
16
  try:
17
+ # Carga segura de artefactos
18
  model = joblib.load(os.path.join(os.path.dirname(__file__), 'model.pkl'))
19
  preprocessor = joblib.load(os.path.join(os.path.dirname(__file__), 'preprocessor.pkl'))
20
  model_columns = joblib.load(os.path.join(os.path.dirname(__file__), 'model_columns.pkl'))
21
  except Exception as e:
22
+ # Este error capturar谩 el problema de compatibilidad de Scikit-learn
23
  print(f"Error al cargar artefactos: {e}")
24
  model = None
25
  preprocessor = None
26
  model_columns = []
27
 
28
  # --- 2. Funci贸n de Predicci贸n (N煤cleo de la API) ---
 
29
  def predict_ko_tko(F1_KD, F2_KD, F1_STR, F2_STR, F1_TD, F2_TD, F1_SUB, F2_SUB, Round,
30
  F1_acc, F2_acc, KD_diff, STR_diff, TD_diff, SUB_diff, Location,
 
31
  wc_B, wc_C, wc_F, wc_Fl, wc_H, wc_LH, wc_L, wc_M, wc_O, wc_SH, wc_W,
32
  wc_WB, wc_WF, wc_WFl, wc_WS):
33
 
34
  if not model or not preprocessor:
35
+ # Devuelve un mensaje claro si la carga fall贸 por la versi贸n de Scikit-learn
36
+ return "ERROR", "Fallo al cargar modelo. Revisa el log de versiones."
37
 
38
+ # 1. Crear el DataFrame de entrada (31 columnas originales de X)
39
  input_data = pd.DataFrame({
40
  'Fighter_1_KD': [F1_KD], 'Fighter_2_KD': [F2_KD], 'Fighter_1_STR': [F1_STR], 'Fighter_2_STR': [F2_STR],
41
  'Fighter_1_TD': [F1_TD], 'Fighter_2_TD': [F2_TD], 'Fighter_1_SUB': [F1_SUB], 'Fighter_2_SUB': [F2_SUB],
 
51
  "weight_class_Women's Flyweight": [wc_WFl], "weight_class_Women's Strawweight": [wc_WS]
52
  })
53
 
54
+ # 2. Preprocesamiento: Utilizar el ColumnTransformer ajustado.
 
55
  X_processed = preprocessor.transform(input_data)
56
 
57
+ # 3. Convertir a DataFrame y asegurar el orden de las columnas
58
  X_final = pd.DataFrame(X_processed, columns=model_columns)
59
 
60
  # 4. Predicci贸n
 
67
  return prob_str, result_str
68
 
69
  # --- 3. Creaci贸n de la Interfaz Gradio ---
 
70
  inputs = [
71
  gr.Slider(0, 5, value=1, step=1, label="KD P1"), gr.Slider(0, 5, value=0, step=1, label="KD P2"),
72
  gr.Slider(0, 300, value=70, label="STR P1"), gr.Slider(0, 300, value=50, label="STR P2"),
 
76
  gr.Slider(0, 1, value=0.4, label="Precisi贸n STR P1 (F1_acc)"), gr.Slider(0, 1, value=0.3, label="Precisi贸n STR P2 (F2_acc)"),
77
  gr.Slider(-5, 5, value=1, label="Diferencia de KD"), gr.Slider(-300, 300, value=20, label="Diferencia de STR"),
78
  gr.Slider(-20, 20, value=3, label="Diferencia de TD"), gr.Slider(-5, 5, value=0, label="Diferencia de SUB"),
79
+ gr.Dropdown(UFC_LOCATIONS, value='Las Vegas, NV', label="Ubicaci贸n"), # Usa la lista corregida
80
+ gr.Checkbox(value=True, label="weight_class_Lightweight"), gr.Checkbox(value=False, label="weight_class_Bantamweight"),
81
+ gr.Checkbox(value=False, label="weight_class_Catch Weight"), gr.Checkbox(value=False, label="weight_class_Featherweight"),
82
+ gr.Checkbox(value=False, label="weight_class_Flyweight"), gr.Checkbox(value=False, label="weight_class_Heavyweight"),
83
+ gr.Checkbox(value=False, label="weight_class_Light Heavyweight"), gr.Checkbox(value=False, label="weight_class_Middleweight"),
84
+ gr.Checkbox(value=False, label="weight_class_Open Weight"), gr.Checkbox(value=False, label="weight_class_Super Heavyweight"),
85
+ gr.Checkbox(value=False, label="weight_class_Welterweight"), gr.Checkbox(value=False, label="weight_class_Women's Bantamweight"),
86
+ gr.Checkbox(value=False, label="weight_class_Women's Featherweight"), gr.Checkbox(value=False, label="weight_class_Women's Flyweight"),
87
+ gr.Checkbox(value=False, label="weight_class_Women's Strawweight")
 
88
  ]
89
 
90
  outputs = [gr.Textbox(label="Probabilidad de KO/TKO (%)"), gr.Textbox(label="Resultado M谩s Probable")]
 
94
  inputs=inputs,
95
  outputs=outputs,
96
  title="馃 Predictor de KO/TKO en Combates UFC (Despliegue ML)",
97
+ description="Modelo Random Forest para predecir la finalizaci贸n de un combate."
98
+ ).launch(server_name="0.0.0.0", server_port=7860)