jersonalvr/machinelearning
0
1# utils/models_utils.py2import streamlit as st3import pandas as pd4import numpy as np5import plotly.express as px6import time7import pickle8import io9from stqdm import stqdm10from sklearn.model_selection import GridSearchCV, train_test_split11from sklearn.pipeline import Pipeline12from sklearn.preprocessing import StandardScaler, LabelEncoder13from sklearn.linear_model import (14 LinearRegression, LogisticRegression, Lasso, Ridge,15 SGDClassifier, RidgeClassifier, PassiveAggressiveClassifier16)17from sklearn.tree import DecisionTreeRegressor, DecisionTreeClassifier18from sklearn.ensemble import (19 RandomForestRegressor, RandomForestClassifier,20 GradientBoostingClassifier, AdaBoostClassifier,21 BaggingClassifier, ExtraTreesClassifier, ExtraTreesRegressor22)23from sklearn.naive_bayes import GaussianNB, MultinomialNB, BernoulliNB24from sklearn.neighbors import KNeighborsClassifier25from sklearn.svm import SVC, SVR26from sklearn.metrics import (27 mean_squared_error, r2_score, mean_absolute_error,28 accuracy_score, classification_report, confusion_matrix29)30from sklearn.base import BaseEstimator, ClassifierMixin, RegressorMixin31import xgboost as xgb32import h2o33import os34 35class ModelTrainer:36 """37 Clase para gestionar el entrenamiento de modelos de machine learning38 """39 @staticmethod40 def get_model_options(problem_type):41 """42 Obtener opciones de modelos según el tipo de problema43 44 Args:45 problem_type (str): Tipo de problema ('classification' o 'regression')46 47 Returns:48 dict: Diccionario de opciones de modelos49 """50 if problem_type == 'regression':51 return ModelTrainer._get_regression_models()52 else:53 return ModelTrainer._get_classification_models()54 55 @staticmethod56 def _get_regression_models():57 """58 Definir opciones de modelos para regresión59 60 Returns:61 dict: Modelos de regresión con sus parámetros62 """63 return {64 'Regresión Lineal': {65 'model': lambda rs: Pipeline([66 ('scaler', StandardScaler()),67 ('regressor', LinearRegression())68 ]),69 'params': {70 'regressor__fit_intercept': [True, False],71 'regressor__copy_X': [True],72 'regressor__positive': [True, False],73 'scaler__with_mean': [True, False],74 'scaler__with_std': [True, False]75 }76 },77 'Lasso': {78 'model': lambda rs: Pipeline([79 ('scaler', StandardScaler()),80 ('regressor', Lasso(random_state=rs))81 ]),82 'params': {83 'regressor__alpha': [0.0001, 0.001, 0.01, 0.1, 1.0, 10.0],84 'regressor__fit_intercept': [True, False],85 'regressor__max_iter': [1000, 2000, 5000],86 'regressor__selection': ['cyclic', 'random'],87 'regressor__tol': [1e-4, 1e-3],88 'scaler__with_mean': [True, False],89 'scaler__with_std': [True, False]90 }91 },92 'Ridge': {93 'model': lambda rs: Pipeline([94 ('scaler', StandardScaler()),95 ('regressor', Ridge(random_state=rs))96 ]),97 'params': {98 'regressor__alpha': [0.0001, 0.001, 0.01, 0.1, 1.0, 10.0],99 'regressor__fit_intercept': [True, False],100 'regressor__solver': ['auto', 'svd', 'cholesky', 'lsqr', 'sparse_cg', 'sag', 'saga'],101 'regressor__tol': [1e-4, 1e-3],102 'scaler__with_mean': [True, False],103 'scaler__with_std': [True, False]104 }105 },106 'Árbol de Decisión': {107 'model': lambda rs: DecisionTreeRegressor(random_state=rs),108 'params': {109 'max_depth': [3, 5, 7, 10, 15, None],110 'min_samples_split': [2, 5, 10, 20],111 'min_samples_leaf': [1, 2, 4, 8],112 'criterion': ['squared_error', 'friedman_mse', 'absolute_error', 'poisson'],113 'splitter': ['best', 'random'],114 'max_features': ['sqrt', 'log2', None]115 }116 },117 'Random Forest': {118 'model': lambda rs: RandomForestRegressor(random_state=rs),119 'params': {120 'n_estimators': [100, 200, 300, 500],121 'max_depth': [3, 5, 7, 10, None],122 'min_samples_split': [2, 5, 10, 20],123 'min_samples_leaf': [1, 2, 4],124 'max_features': ['sqrt', 'log2', None],125 'bootstrap': [True, False],126 'criterion': ['squared_error', 'absolute_error', 'poisson']127 }128 },129 'XGBoost': {130 'model': lambda rs: xgb.XGBRegressor(131 tree_method='hist',132 device='cuda',133 enable_categorical=True,134 random_state=rs135 ),136 'params': {137 'n_estimators': [100, 200, 300, 500],138 'max_depth': [3, 5, 7, 9],139 'learning_rate': [0.01, 0.05, 0.1, 0.3],140 'subsample': [0.8, 0.9, 1.0],141 'colsample_bytree': [0.8, 0.9, 1.0],142 'min_child_weight': [1, 3, 5],143 'gamma': [0, 0.1, 0.2],144 'reg_alpha': [0, 0.1, 0.5],145 'reg_lambda': [0.1, 1.0, 5.0]146 }147 }148 }149 150 @staticmethod151 def _get_classification_models():152 """153 Definir opciones de modelos para clasificación154 155 Returns:156 dict: Modelos de clasificación con sus parámetros157 """158 return {159 'Regresión Logística': {160 'model': lambda rs: LogisticRegression(max_iter=1000, random_state=rs),161 'params': {162 'C': [0.001, 0.01, 0.1, 1.0, 10.0],163 'penalty': ['l1', 'l2'],164 'solver': ['liblinear', 'saga'],165 'class_weight': [None, 'balanced'],166 'warm_start': [True, False],167 'tol': [1e-4, 1e-3, 1e-2]168 }169 },170 'Random Forest': {171 'model': lambda rs: RandomForestClassifier(random_state=rs),172 'params': {173 'n_estimators': [100, 200, 300, 500],174 'max_depth': [3, 5, 7, 10, None],175 'min_samples_split': [2, 5, 10],176 'min_samples_leaf': [1, 2, 4],177 'class_weight': [None, 'balanced', 'balanced_subsample'],178 'criterion': ['gini', 'entropy'],179 'max_features': ['sqrt', 'log2', None]180 }181 },182 'XGBoost': {183 'model': lambda rs: xgb.XGBClassifier(184 tree_method='hist',185 device='cuda',186 enable_categorical=True,187 random_state=rs188 ),189 'params': {190 'n_estimators': [100, 200, 300, 500],191 'max_depth': [3, 5, 7, 9],192 'learning_rate': [0.01, 0.05, 0.1, 0.3],193 'subsample': [0.8, 0.9, 1.0],194 'colsample_bytree': [0.8, 0.9, 1.0],195 'min_child_weight': [1, 3, 5],196 'gamma': [0, 0.1, 0.2],197 'reg_alpha': [0, 0.1, 0.5],198 'reg_lambda': [0.1, 1.0, 5.0],199 'scale_pos_weight': [1, 2, 3]200 }201 },202 'SVM': {203 'model': lambda rs: SVC(random_state=rs),204 'params': {205 'C': [0.1, 1, 10, 100],206 'kernel': ['linear', 'rbf', 'poly', 'sigmoid'],207 'gamma': ['scale', 'auto', 0.1, 0.01, 0.001],208 'class_weight': [None, 'balanced'],209 'probability': [True]210 }211 },212 'Naive Bayes': {213 'model': lambda rs: GaussianNB(),214 'params': {215 'var_smoothing': [1e-9, 1e-8, 1e-7, 1e-6]216 }217 }218 }219 220 @staticmethod221 def _determine_problem_type(model):222 """223 Determinar el tipo de problema basado en el modelo224 225 Args:226 model (BaseEstimator): Modelo a evaluar227 228 Returns:229 str: Tipo de problema ('classification', 'regression', 'unknown')230 """231 try:232 if hasattr(model, 'predict_proba'):233 return 'classification'234 elif hasattr(model, 'predict'):235 return 'regression'236 else:237 return 'unknown'238 except ImportError:239 return 'unknown'240 241 @staticmethod242 def _get_default_scoring(problem_type):243 """244 Obtener la métrica de scoring predeterminada245 246 Args:247 problem_type (str): Tipo de problema248 249 Returns:250 str: Métrica de scoring predeterminada251 """252 scoring_map = {253 'classification': 'accuracy',254 'regression': 'r2',255 'unknown': None256 }257 return scoring_map.get(problem_type, None)258 259 @staticmethod260 def train_model_pipeline(261 X_train, 262 y_train, 263 model_config, 264 X_test=None, 265 y_test=None, 266 cv=5, 267 scoring=None, 268 random_state=42, 269 **kwargs270 ):271 """272 Entrenar modelo con validación cruzada y evaluación flexible273 274 Args:275 X_train (array-like): Datos de entrenamiento276 y_train (array-like): Etiquetas de entrenamiento277 model_config (dict): Configuración del modelo278 X_test (array-like, optional): Datos de prueba279 y_test (array-like, optional): Etiquetas de prueba280 cv (int, optional): Número de pliegues para validación cruzada281 scoring (str, optional): Métrica de puntuación282 random_state (int, optional): Semilla aleatoria para reproducibilidad283 **kwargs: Argumentos adicionales284 285 Returns:286 dict: Resultados detallados del entrenamiento287 """288 # Extraer modelo y parámetros289 model_func = model_config.get('model')290 params = model_config.get('params', {})291 292 # Instanciar el modelo si es una función293 if callable(model_func):294 model = model_func(random_state)295 else:296 model = model_func297 298 # Verificar que el modelo sea una instancia válida299 if not hasattr(model, 'fit') or not hasattr(model, 'predict'):300 raise ValueError(f"Modelo inválido: {model}. Debe tener métodos 'fit' y 'predict'.")301 302 # Determinar tipo de problema303 problem_type = ModelTrainer._determine_problem_type(model)304 305 # Configurar scoring306 if scoring is None:307 scoring = ModelTrainer._get_default_scoring(problem_type)308 309 # Configurar parámetros de GridSearchCV310 grid_search_params = {311 'estimator': model,312 'param_grid': params,313 'cv': cv,314 'scoring': scoring315 }316 317 # Añadir kwargs adicionales318 grid_search_params.update({319 k: v for k, v in kwargs.items() 320 if k in ['n_jobs', 'verbose', 'refit', 'error_score']321 })322 323 try:324 # Realizar búsqueda de hiperparámetros325 grid_search = GridSearchCV(**grid_search_params)326 with st.spinner(f"Entrenando modelo {model}..."):327 start_time = time.time()328 grid_search.fit(X_train, y_train)329 training_time = time.time() - start_time330 331 except Exception as e:332 return {333 'error': f"Error durante el entrenamiento: {str(e)}",334 'problem_type': problem_type335 }336 337 # Preparar resultados base338 results = {339 'problem_type': problem_type,340 'best_model': grid_search.best_estimator_,341 'best_params': grid_search.best_params_,342 'best_score': grid_search.best_score_,343 'cv_results': grid_search.cv_results_,344 'training_time': training_time345 }346 347 # Evaluación en conjunto de prueba348 if X_test is not None and y_test is not None:349 best_model = grid_search.best_estimator_350 y_pred = best_model.predict(X_test)351 352 # Métricas específicas según el tipo de problema353 if problem_type == 'classification':354 results.update({355 'test_accuracy': accuracy_score(y_test, y_pred),356 'classification_report': classification_report(y_test, y_pred, output_dict=True),357 'confusion_matrix': confusion_matrix(y_test, y_pred).tolist(),358 'y_pred': y_pred359 })360 elif problem_type == 'regression':361 results.update({362 'test_mse': mean_squared_error(y_test, y_pred),363 'test_rmse': np.sqrt(mean_squared_error(y_test, y_pred)),364 'test_mae': mean_absolute_error(y_test, y_pred),365 'test_r2': r2_score(y_test, y_pred),366 'y_pred': y_pred367 })368 else:369 results['test_predictions'] = y_pred370 371 return results372 373 @staticmethod374 def create_class_distribution_plot(y_original):375 """376 Crear un gráfico de distribución de clases377 378 Args:379 y_original (pd.Series): Variable objetivo original380 381 Returns:382 plotly.graph_objs._figure.Figure: Gráfico de distribución de clases383 """384 class_dist = pd.DataFrame({385 'Clase': y_original.value_counts().index,386 'Cantidad': y_original.value_counts().values387 })388 389 fig = px.bar(390 class_dist,391 x='Clase',392 y='Cantidad',393 title='Distribución de clases'394 )395 396 return fig397 398 @staticmethod399 def process_classification_data(y, random_state):400 """401 Procesar datos de clasificación402 403 Args:404 y (pd.Series): Variable objetivo405 random_state (int): Semilla aleatoria406 407 Returns:408 tuple: Variable objetivo procesada y codificador de etiquetas409 """410 # Codificación de etiquetas411 le = LabelEncoder()412 y_encoded = pd.Series(le.fit_transform(y))413 414 return y_encoded, le415 416 @staticmethod417 def save_model(model, filename):418 """419 Guardar modelo entrenado en un archivo420 421 Args:422 model: Modelo entrenado423 filename (str): Nombre del archivo424 """425 if isinstance(model, h2o.estimators.H2OEstimator):426 # Usar método nativo de H2O para guardar modelos427 h2o.save_model(model=model, path=os.path.dirname(filename), force=True)428 else:429 with open(filename, 'wb') as f:430 pickle.dump(model, f)431 432 @staticmethod433 def load_model(filename):434 """435 Cargar modelo desde un archivo436 437 Args:438 filename (str): Nombre del archivo439 440 Returns:441 Modelo cargado442 """443 if filename.endswith('.zip'):444 # Asumir que es un modelo H2O445 return h2o.load_model(filename)446 else:447 with open(filename, 'rb') as f:448 return pickle.load(f)449 450 @staticmethod451 def get_model_performance_metrics(y_true, y_pred, problem_type):452 """453 Obtener métricas de rendimiento del modelo454 455 Args:456 y_true (pd.Series): Etiquetas verdaderas457 y_pred (pd.Series): Etiquetas predichas458 problem_type (str): Tipo de problema459 460 Returns:461 dict: Métricas de rendimiento462 """463 if problem_type == 'classification':464 return {465 'accuracy': accuracy_score(y_true, y_pred),466 'classification_report': classification_report(y_true, y_pred, output_dict=True)467 }468 else: # Regresión469 return {470 'mse': mean_squared_error(y_true, y_pred),471 'r2_score': r2_score(y_true, y_pred)472 }473 474 @staticmethod475 def split_data(X, y, test_size=0.2, random_state=42):476 """477 Dividir datos en conjuntos de entrenamiento y prueba478 479 Args:480 X (pd.DataFrame): Features481 y (pd.Series): Variable objetivo482 test_size (float): Proporción de datos de prueba483 random_state (int): Semilla aleatoria484 485 Returns:486 tuple: X_train, X_test, y_train, y_test487 """488 return train_test_split(X, y, test_size=test_size, random_state=random_state)489 490 @staticmethod491 def prepare_data_for_ml(df, target_column, problem_type='classification', test_size=0.2, random_state=42):492 """493 Preparar datos para machine learning494 495 Args:496 df (pd.DataFrame): DataFrame de datos497 target_column (str): Columna objetivo498 problem_type (str): Tipo de problema499 test_size (float): Proporción de datos de prueba500 random_state (int): Semilla aleatoria501 502 Returns:503 dict: Diccionario con datos preparados504 """505 # Separar features y target506 X = df.drop(columns=[target_column])507 y = df[target_column]508 509 # Preprocesar datos según el tipo de problema510 if problem_type == 'classification':511 y, label_encoder = ModelTrainer.process_classification_data(y, random_state)512 else:513 label_encoder = None514 515 # Dividir datos516 X_train, X_test, y_train, y_test = ModelTrainer.split_data(X, y, test_size, random_state)517 518 return {519 'X_train': X_train,520 'X_test': X_test,521 'y_train': y_train,522 'y_test': y_test,523 'label_encoder': label_encoder,524 'features': list(X.columns),525 'problem_type': problem_type526 }527 528 @staticmethod529 def generate_model_comparison_report(trained_models, problem_type):530 """531 Generar informe comparativo de modelos532 533 Args:534 trained_models (dict): Modelos entrenados535 problem_type (str): Tipo de problema536 537 Returns:538 pd.DataFrame: Informe comparativo de modelos539 """540 comparison_data = []541 542 for model_name, model_info in trained_models.items():543 model_metrics = ModelTrainer.get_model_performance_metrics(544 model_info['y_test'], 545 model_info['y_pred'], 546 problem_type547 )548 549 model_entry = {550 'Modelo': model_name,551 'Tiempo de Entrenamiento': model_info.get('training_time', 0),552 }553 554 # Agregar métricas según el tipo de problema555 if problem_type == 'classification':556 model_entry.update({557 'Precisión': model_metrics['accuracy'],558 'Precisión (Macro)': model_metrics['classification_report']['macro avg']['precision'],559 'Recall (Macro)': model_metrics['classification_report']['macro avg']['recall'],560 'F1-Score (Macro)': model_metrics['classification_report']['macro avg']['f1-score']561 })562 else:563 model_entry.update({564 'MSE': model_metrics['mse'],565 'R2 Score': model_metrics['r2_score']566 })567 568 comparison_data.append(model_entry)569 570 return pd.DataFrame(comparison_data)571 572 @staticmethod573 def plot_model_comparison(comparison_df, problem_type):574 """575 Crear gráfico comparativo de modelos576 577 Args:578 comparison_df (pd.DataFrame): DataFrame de comparación de modelos579 problem_type (str): Tipo de problema580 581 Returns:582 plotly.graph_objs._figure.Figure: Gráfico comparativo583 """584 metric_column = 'Precisión' if problem_type == 'classification' else 'R2 Score'585 586 fig = px.bar(587 comparison_df, 588 x='Modelo', 589 y=metric_column,590 title=f'Comparación de Modelos - {metric_column}'591 )592 593 return fig594 595# Funciones sueltas para importación directa596def get_model_options(problem_type):597 return ModelTrainer.get_model_options(problem_type)598 599def train_model_pipeline(*args, **kwargs):600 return ModelTrainer.train_model_pipeline(*args, **kwargs)601 602def process_classification_data(y, random_state=42):603 return ModelTrainer.process_classification_data(y, random_state)604 605def create_class_distribution_plot(y):606 return ModelTrainer.create_class_distribution_plot(y)