Ajuste de los hiperparámetros de un estimador
Índice de contenido
¿Qué son los hiperparámetros?
Los hiperparámetros son parámetros que no se aprenden directamente dentro de los estimadores. En scikit-learn, se pasan como argumentos al constructor de las clases. Hiperparámetros típicos son, por ejemplo, “C”, “kernel” y “gamma” para los clasificadores de vectores de soporte.
¿Se pueden optimizar los hiperparámetros?
Es posible y recomendable buscar en el espacio de hiperparámetros los mejores valores para, a su vez, obtener la mejor puntuación de validación cruzada. Cualquier parámetro proporcionado al construir un estimador puede optimizarse de esta manera.
¿Qué hace falta para llevar a cabo una búsqueda de hiperparámetros?
Una búsqueda consta de:
- Un estimador (un objeto regresor o clasificador, como por ejemplo sklearn.svm.SVC())
- Un espacio de parámetros en el cual buscar
- Un método para buscar o muestrear candidatos
- Un esquema de validación cruzada
- Una función de puntuación
Métodos de búsqueda de hiperparámetros
En Scikit-Learn se proporcionan dos enfoques genéricos para buscar hiperparámetros:
- Para un conjunto de valores dados, GridSearchCV considera exhaustivamente todas las combinaciones posibles de dichos valores de hiperparámetros.
- El método RandomizedSearchCV puede analizar un número determinado de valores de un espacio de hiperparámetros con una distribución específica.
Es importante tener en cuenta que puede haber unos pocos hiperparámetros que tengan un gran impacto en el rendimiento del modelo, mientras que otros quizás puedan dejarse simplemente con los valores que les da Scikit-Learn por defecto.
GridSearchCV (búsqueda exhaustiva de cuadrícula)
La búsqueda de cuadrícula (GridSearchCV en Scikit-Learn) estudia todas las posibles combinaciones de valores de hiperparámetros, que se especifican mediante el parámetro param_grid, como por ejemplo de la siguiente manera:
param_grid = [
{‘C’: [1, 10, 100, 1000], ‘kernel’: [’linear’]},
{‘C’: [1, 10, 100, 1000], ‘gamma’: [0.001, 0.0001], ‘kernel’: [‘rbf’]},
]
Este ejemplo de param_grid indica que se deben explorar dos cuadrículas: una con un kernel lineal y valores C en [1, 10, 100, 1000], y la segunda con un kernel RBF, y el producto cruzado de los valores C en [1, 10 , 100, 1000] y valores gamma en [0.001, 0.0001].
Tras evaluar todas las combinaciones, se conserva la mejor de ellas.
¿Cómo se usa GridSearchCV en Scikit-Learn?
class sklearn.model_selection.GridSearchCV(estimator, param_grid, , scoring=None, n_jobs=None, iid=‘deprecated’, refit=True, cv=None, verbose=0, pre_dispatch=‘2n_jobs’, error_score=nan, return_train_score=False)
Veamos los hiperparámetros más importantes de esta clase:
- estimator: un estimador Scikit-Learn, como por ejemplo sklearn.svm.SVC().
- param_grid: diccionario o lista de diccionarios. Debe proporcionar los nombres de los hiperparámetros con los valores que se desean probar.
- scoring: métrica/s para evaluar las predicciones en el conjunto de prueba. Si es None, se utiliza la métrica de puntuación que tenga asociada el estimador.
- refit: indica si se desea reajustar el estimador utilizando los mejores parámetros encontrados en todo el conjunto de datos. Ese estimador reajustado estará disponible en el atributo best_estimator_ y se podrá utilizar para hacer predicciones directamente con la instancia de GridSearchCV.
- cv: determina la estrategia de división de validación cruzada. Consultar las posibles entradas para cv en la documentación oficial.
- verbose: controla la verbosidad (información que se muestra durante la ejecución). Se indica mediante un número entero, y cuanto más alto sea su valor, más mensajes se mostrarán.
- error_score: valor para asignar a la puntuación si se produce un error en el ajuste del estimador.
- return_train_score: si es False, el atributo cv_results_ no incluirá las puntuaciones de entrenamiento. El cálculo de puntuaciones de entrenamiento se usa para obtener información sobre cómo diferentes configuraciones de parámetros impactan en la compensación de overfitting/underfitting. Sin embargo, calcular los resultados de las métricas de evaluación en el conjunto de entrenamiento puede ser computacionalmente costoso y no es estrictamente necesario para seleccionar los parámetros que producen el mejor rendimiento de generalización.
Ejemplo práctico
from sklearn import datasets
from sklearn import svm
from sklearn.model_selection import GridSearchCV
import pandas as pd
iris = datasets.load_iris()
iris
{'data': array([[5.1, 3.5, 1.4, 0.2],
[4.9, 3. , 1.4, 0.2],
[4.7, 3.2, 1.3, 0.2],
[4.6, 3.1, 1.5, 0.2],
[5. , 3.6, 1.4, 0.2],
[5.4, 3.9, 1.7, 0.4],
[4.6, 3.4, 1.4, 0.3],
[5. , 3.4, 1.5, 0.2],
[4.4, 2.9, 1.4, 0.2],
[4.9, 3.1, 1.5, 0.1],
[5.4, 3.7, 1.5, 0.2],
[4.8, 3.4, 1.6, 0.2],
[4.8, 3. , 1.4, 0.1],
[4.3, 3. , 1.1, 0.1],
[5.8, 4. , 1.2, 0.2],
[5.7, 4.4, 1.5, 0.4],
[5.4, 3.9, 1.3, 0.4],
[5.1, 3.5, 1.4, 0.3],
[5.7, 3.8, 1.7, 0.3],
[5.1, 3.8, 1.5, 0.3],
[5.4, 3.4, 1.7, 0.2],
[5.1, 3.7, 1.5, 0.4],
[4.6, 3.6, 1. , 0.2],
[5.1, 3.3, 1.7, 0.5],
[4.8, 3.4, 1.9, 0.2],
[5. , 3. , 1.6, 0.2],
[5. , 3.4, 1.6, 0.4],
[5.2, 3.5, 1.5, 0.2],
[5.2, 3.4, 1.4, 0.2],
[4.7, 3.2, 1.6, 0.2],
[4.8, 3.1, 1.6, 0.2],
[5.4, 3.4, 1.5, 0.4],
[5.2, 4.1, 1.5, 0.1],
[5.5, 4.2, 1.4, 0.2],
[4.9, 3.1, 1.5, 0.2],
[5. , 3.2, 1.2, 0.2],
[5.5, 3.5, 1.3, 0.2],
[4.9, 3.6, 1.4, 0.1],
[4.4, 3. , 1.3, 0.2],
[5.1, 3.4, 1.5, 0.2],
[5. , 3.5, 1.3, 0.3],
[4.5, 2.3, 1.3, 0.3],
[4.4, 3.2, 1.3, 0.2],
[5. , 3.5, 1.6, 0.6],
[5.1, 3.8, 1.9, 0.4],
[4.8, 3. , 1.4, 0.3],
[5.1, 3.8, 1.6, 0.2],
[4.6, 3.2, 1.4, 0.2],
[5.3, 3.7, 1.5, 0.2],
[5. , 3.3, 1.4, 0.2],
[7. , 3.2, 4.7, 1.4],
[6.4, 3.2, 4.5, 1.5],
[6.9, 3.1, 4.9, 1.5],
[5.5, 2.3, 4. , 1.3],
[6.5, 2.8, 4.6, 1.5],
[5.7, 2.8, 4.5, 1.3],
[6.3, 3.3, 4.7, 1.6],
[4.9, 2.4, 3.3, 1. ],
[6.6, 2.9, 4.6, 1.3],
[5.2, 2.7, 3.9, 1.4],
[5. , 2. , 3.5, 1. ],
[5.9, 3. , 4.2, 1.5],
[6. , 2.2, 4. , 1. ],
[6.1, 2.9, 4.7, 1.4],
[5.6, 2.9, 3.6, 1.3],
[6.7, 3.1, 4.4, 1.4],
[5.6, 3. , 4.5, 1.5],
[5.8, 2.7, 4.1, 1. ],
[6.2, 2.2, 4.5, 1.5],
[5.6, 2.5, 3.9, 1.1],
[5.9, 3.2, 4.8, 1.8],
[6.1, 2.8, 4. , 1.3],
[6.3, 2.5, 4.9, 1.5],
[6.1, 2.8, 4.7, 1.2],
[6.4, 2.9, 4.3, 1.3],
[6.6, 3. , 4.4, 1.4],
[6.8, 2.8, 4.8, 1.4],
[6.7, 3. , 5. , 1.7],
[6. , 2.9, 4.5, 1.5],
[5.7, 2.6, 3.5, 1. ],
[5.5, 2.4, 3.8, 1.1],
[5.5, 2.4, 3.7, 1. ],
[5.8, 2.7, 3.9, 1.2],
[6. , 2.7, 5.1, 1.6],
[5.4, 3. , 4.5, 1.5],
[6. , 3.4, 4.5, 1.6],
[6.7, 3.1, 4.7, 1.5],
[6.3, 2.3, 4.4, 1.3],
[5.6, 3. , 4.1, 1.3],
[5.5, 2.5, 4. , 1.3],
[5.5, 2.6, 4.4, 1.2],
[6.1, 3. , 4.6, 1.4],
[5.8, 2.6, 4. , 1.2],
[5. , 2.3, 3.3, 1. ],
[5.6, 2.7, 4.2, 1.3],
[5.7, 3. , 4.2, 1.2],
[5.7, 2.9, 4.2, 1.3],
[6.2, 2.9, 4.3, 1.3],
[5.1, 2.5, 3. , 1.1],
[5.7, 2.8, 4.1, 1.3],
[6.3, 3.3, 6. , 2.5],
[5.8, 2.7, 5.1, 1.9],
[7.1, 3. , 5.9, 2.1],
[6.3, 2.9, 5.6, 1.8],
[6.5, 3. , 5.8, 2.2],
[7.6, 3. , 6.6, 2.1],
[4.9, 2.5, 4.5, 1.7],
[7.3, 2.9, 6.3, 1.8],
[6.7, 2.5, 5.8, 1.8],
[7.2, 3.6, 6.1, 2.5],
[6.5, 3.2, 5.1, 2. ],
[6.4, 2.7, 5.3, 1.9],
[6.8, 3. , 5.5, 2.1],
[5.7, 2.5, 5. , 2. ],
[5.8, 2.8, 5.1, 2.4],
[6.4, 3.2, 5.3, 2.3],
[6.5, 3. , 5.5, 1.8],
[7.7, 3.8, 6.7, 2.2],
[7.7, 2.6, 6.9, 2.3],
[6. , 2.2, 5. , 1.5],
[6.9, 3.2, 5.7, 2.3],
[5.6, 2.8, 4.9, 2. ],
[7.7, 2.8, 6.7, 2. ],
[6.3, 2.7, 4.9, 1.8],
[6.7, 3.3, 5.7, 2.1],
[7.2, 3.2, 6. , 1.8],
[6.2, 2.8, 4.8, 1.8],
[6.1, 3. , 4.9, 1.8],
[6.4, 2.8, 5.6, 2.1],
[7.2, 3. , 5.8, 1.6],
[7.4, 2.8, 6.1, 1.9],
[7.9, 3.8, 6.4, 2. ],
[6.4, 2.8, 5.6, 2.2],
[6.3, 2.8, 5.1, 1.5],
[6.1, 2.6, 5.6, 1.4],
[7.7, 3. , 6.1, 2.3],
[6.3, 3.4, 5.6, 2.4],
[6.4, 3.1, 5.5, 1.8],
[6. , 3. , 4.8, 1.8],
[6.9, 3.1, 5.4, 2.1],
[6.7, 3.1, 5.6, 2.4],
[6.9, 3.1, 5.1, 2.3],
[5.8, 2.7, 5.1, 1.9],
[6.8, 3.2, 5.9, 2.3],
[6.7, 3.3, 5.7, 2.5],
[6.7, 3. , 5.2, 2.3],
[6.3, 2.5, 5. , 1.9],
[6.5, 3. , 5.2, 2. ],
[6.2, 3.4, 5.4, 2.3],
[5.9, 3. , 5.1, 1.8]]),
'target': array([0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0,
0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0,
0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1,
1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1,
1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2,
2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2,
2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2]),
'frame': None,
'target_names': array(['setosa', 'versicolor', 'virginica'], dtype='<U10'),
'DESCR': '.. _iris_dataset:\n\nIris plants dataset\n--------------------\n\n**Data Set Characteristics:**\n\n :Number of Instances: 150 (50 in each of three classes)\n :Number of Attributes: 4 numeric, predictive attributes and the class\n :Attribute Information:\n - sepal length in cm\n - sepal width in cm\n - petal length in cm\n - petal width in cm\n - class:\n - Iris-Setosa\n - Iris-Versicolour\n - Iris-Virginica\n \n :Summary Statistics:\n\n ============== ==== ==== ======= ===== ====================\n Min Max Mean SD Class Correlation\n ============== ==== ==== ======= ===== ====================\n sepal length: 4.3 7.9 5.84 0.83 0.7826\n sepal width: 2.0 4.4 3.05 0.43 -0.4194\n petal length: 1.0 6.9 3.76 1.76 0.9490 (high!)\n petal width: 0.1 2.5 1.20 0.76 0.9565 (high!)\n ============== ==== ==== ======= ===== ====================\n\n :Missing Attribute Values: None\n :Class Distribution: 33.3% for each of 3 classes.\n :Creator: R.A. Fisher\n :Donor: Michael Marshall (MARSHALL%PLU@io.arc.nasa.gov)\n :Date: July, 1988\n\nThe famous Iris database, first used by Sir R.A. Fisher. The dataset is taken\nfrom Fisher\'s paper. Note that it\'s the same as in R, but not as in the UCI\nMachine Learning Repository, which has two wrong data points.\n\nThis is perhaps the best known database to be found in the\npattern recognition literature. Fisher\'s paper is a classic in the field and\nis referenced frequently to this day. (See Duda & Hart, for example.) The\ndata set contains 3 classes of 50 instances each, where each class refers to a\ntype of iris plant. One class is linearly separable from the other 2; the\nlatter are NOT linearly separable from each other.\n\n.. topic:: References\n\n - Fisher, R.A. "The use of multiple measurements in taxonomic problems"\n Annual Eugenics, 7, Part II, 179-188 (1936); also in "Contributions to\n Mathematical Statistics" (John Wiley, NY, 1950).\n - Duda, R.O., & Hart, P.E. (1973) Pattern Classification and Scene Analysis.\n (Q327.D83) John Wiley & Sons. ISBN 0-471-22361-1. See page 218.\n - Dasarathy, B.V. (1980) "Nosing Around the Neighborhood: A New System\n Structure and Classification Rule for Recognition in Partially Exposed\n Environments". IEEE Transactions on Pattern Analysis and Machine\n Intelligence, Vol. PAMI-2, No. 1, 67-71.\n - Gates, G.W. (1972) "The Reduced Nearest Neighbor Rule". IEEE Transactions\n on Information Theory, May 1972, 431-433.\n - See also: 1988 MLC Proceedings, 54-64. Cheeseman et al"s AUTOCLASS II\n conceptual clustering system finds 3 classes in the data.\n - Many, many more ...',
'feature_names': ['sepal length (cm)',
'sepal width (cm)',
'petal length (cm)',
'petal width (cm)'],
'filename': 'C:\\Users\\user\\anaconda3\\envs\\python-385\\lib\\site-packages\\sklearn\\datasets\\data\\iris.csv'}
svc = svm.SVC()
paramgrid = {"kernel": ("linear", "rbf"),
"C": [1, 10]}
clf = GridSearchCV(estimator = svc,
param_grid = paramgrid,
scoring=None,
cv=None)
clf.fit(iris.data, iris.target)
GridSearchCV(estimator=SVC(),
param_grid={'C': [1, 10], 'kernel': ('linear', 'rbf')})
pd.DataFrame(clf.cv_results_)
| mean_fit_time | std_fit_time | mean_score_time | std_score_time | param_C | param_kernel | params | split0_test_score | split1_test_score | split2_test_score | split3_test_score | split4_test_score | mean_test_score | std_test_score | rank_test_score | |
|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|
| 0 | 0.0020 | 1.784161e-07 | 0.0018 | 0.001600 | 1 | linear | {'C': 1, 'kernel': 'linear'} | 0.966667 | 1.000000 | 0.966667 | 0.966667 | 1.0 | 0.980000 | 0.016330 | 1 |
| 1 | 0.0024 | 7.999420e-04 | 0.0010 | 0.000632 | 1 | rbf | {'C': 1, 'kernel': 'rbf'} | 0.966667 | 0.966667 | 0.966667 | 0.933333 | 1.0 | 0.966667 | 0.021082 | 4 |
| 2 | 0.0008 | 4.000187e-04 | 0.0006 | 0.000490 | 10 | linear | {'C': 10, 'kernel': 'linear'} | 1.000000 | 1.000000 | 0.900000 | 0.966667 | 1.0 | 0.973333 | 0.038873 | 3 |
| 3 | 0.0010 | 0.000000e+00 | 0.0000 | 0.000000 | 10 | rbf | {'C': 10, 'kernel': 'rbf'} | 0.966667 | 1.000000 | 0.966667 | 0.966667 | 1.0 | 0.980000 | 0.016330 | 1 |
clf.best_estimator_
SVC(C=1, kernel='linear')
clf.best_score_
0.9800000000000001
clf.best_params_
{'C': 1, 'kernel': 'linear'}
clf.scorer_
<function sklearn.metrics._scorer._passthrough_scorer(estimator, *args, **kwargs)>
clf.n_splits_
5
RandomizedSearchCV (búsqueda aleatoria)
Si bien la búsqueda de cuadrícula es el método actualmente más utilizado para la optimización de parámetros, otras técnicas tienen propiedades que también las convierten en buenas opciones. RandomizedSearchCV implementa una búsqueda aleatoria de parámetros, donde cada configuración se muestrea a partir de una distribución de posibles valores de parámetros. Esto tiene dos ventajas principales sobre una búsqueda exhaustiva:
- Se puede elegir un presupuesto independientemente del número de parámetros y valores posibles.
- Agregar parámetros que no influyen en el rendimiento no disminuye la eficiencia.
La especificación de cómo se deben muestrear los parámetros se realiza mediante un diccionario, muy similar a la especificación de parámetros para GridSearchCV. Además, un presupuesto de cálculo, que es el número de candidatos muestreados o iteraciones de muestreo, se especifica mediante el parámetro “n_iter”. Para cada parámetro, se puede especificar una distribución sobre los valores posibles o una lista de opciones discretas (que se muestrearán de manera uniforme):
{‘C’: scipy.stats.expon(scale=100), ‘gamma’: scipy.stats.expon(scale=.1),
‘kernel’: [‘rbf’], ‘class_weight’:[‘balanced’, None]}
Este ejemplo utiliza el módulo scipy.stats, que contiene muchas distribuciones útiles para los parámetros de muestreo, tales como expon, gamma, uniform o randint.
Ampliación
Especificar una métrica objetiva
De forma predeterminada, la búsqueda de parámetros utiliza la función score del estimador para evaluar la configuración de un parámetro, concretamente sklearn.metrics.accuracy_score para clasificación y sklearn.metrics.r2_score para regresión. Para algunas aplicaciones, otras funciones de puntuación son más adecuadas (por ejemplo, en la clasificación no balanceada, la puntuación de precisión no suele ser informativa). Se puede especificar una función de puntuación alternativa a través del parámetro “scoring” de GridSearchCV, RandomizedSearchCV y muchas de las herramientas de validación cruzada.
Especificar múltiples métricas para evaluación
GridSearchCV y RandomizedSearchCV permiten especificar varias métricas para el parámetro “scoring”.
La puntuación multimétrica puede especificarse de dos maneras:
- Como una lista de cadenas de nombres de puntuaciones.
scoring = [‘accuracy’, ‘precision’]
- Como un diccionario que correlaciona el nombre de la puntuación con la función que permite calcular dicha puntuación.
from sklearn.metrics import accuracy_score
from sklearn.metrics import make_scorer
scoring = {‘accuracy’: make_scorer(accuracy_score), ‘prec’: ‘precision’}
Selección de modelos
A menudo, se utilizan los métodos de evaluación de hiperparámetros anteriores para comparar el rendimiento de diferentes modelos entre sí, o de diferentes versiones de un mismo modelo.
Al evaluar cada modelo resultante, es importante hacerlo con datos que no se hayan “visto” durante el proceso de búsqueda de la cuadrícula, es decir, se recomienda dividir los datos en un conjunto de desarrollo (que será el que utilice el método GridSearchCV) y un conjunto de evaluación para calcular métricas de rendimiento.
