Team Ai
Apppublic

jersonalvr/machinelearning

sourceHugging Faceupdated 11mo agoView on Hugging Face
0likes
model_utils.py606 linesDownload Raw Back to utils
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)