ÎncepețiÎncepe gratuit

RandomSearchCV în Scikit Learn

Să exersăm construirea unui obiect RandomizedSearchCV folosind Scikit Learn.

Grila de hiperparametri trebuie să includă max_depth (toate valorile între 5 și 25, inclusiv) și max_features ('auto' și 'sqrt').

Opțiunile dorite pentru obiectul RandomizedSearchCV sunt:

  • Un estimator RandomForestClassifier cu n_estimators egal cu 80.
  • Validare încrucișată cu 3 folduri (cv)
  • Folosește roc_auc pentru a evalua modelele
  • Folosește 4 nuclee pentru procesare paralelă (n_jobs)
  • Asigură-te că reantrenezi cel mai bun model și returnezi scorurile de antrenament
  • Eșantionează doar 5 modele pentru eficiență (n_iter)

Seturile de date X_train și y_train sunt deja încărcate.

Reține că hiperparametrii aleși se găsesc în cv_results_, cu câte o coloană pentru fiecare hiperparametru. De exemplu, coloana pentru hiperparametrul criterion ar fi param_criterion.

Acest exercițiu face parte din cursul

Ajustarea hiperparametrilor în Python

Vezi cursul

Instrucțiuni pentru exercițiu

  • Creează o grilă de hiperparametri conform specificațiilor din contextul de mai sus.
  • Creează un obiect RandomizedSearchCV conform descrierii din contextul de mai sus.
  • Antrenează obiectul RandomizedSearchCV pe datele de antrenament.
  • Indexează în obiectul cv_results_ pentru a afișa valorile alese de procesul de modelare pentru ambii hiperparametri (max_depth și max_features).

Exercițiu interactiv practic

Încearcă acest exercițiu completând acest cod de exemplu.

# Create the parameter grid
param_grid = {'max_depth': list(range(____,26)), 'max_features': [____ , ____]} 

# Create a random search object
random_rf_class = RandomizedSearchCV(
    estimator = ____(n_estimators=____),
    param_distributions = ____, n_iter = ____,
    scoring=____, n_jobs=____, cv = ____, refit=____, return_train_score = ____ )

# Fit to the training data
____.fit(X_train, y_train)

# Print the values used for both hyperparameters
print(random_rf_class.cv_results_[____])
print(random_rf_class.cv_results_[____])
Editează și rulează codul