LoslegenKostenlos starten

Randomisierte Suche

# Call GridSearchCV
grid_search = GridSearchCV(clf, param_grid)

# Fit the model
grid_search.fit(X, y)

Im obigen Codeausschnitt aus der vorherigen Übung ist dir vielleicht aufgefallen, dass die erste Zeile sehr schnell ausgeführt wurde, während der Aufruf von .fit() mehrere Sekunden dauerte.

Das liegt daran, dass .fit() die eigentliche Grid-Suche durchführt – und in unserem Fall war es ein Grid mit vielen unterschiedlichen Kombinationen. Je größer das Hyperparameter-Grid, desto langsamer die Grid-Suche. Um dieses Problem zu lösen, können wir statt wirklich alle Kombinationen auszuprobieren zufällig durch das Grid springen und verschiedene Kombinationen testen. Es besteht eine kleine Chance, dass wir die beste Kombination verpassen, aber wir sparen viel Zeit oder können in derselben Zeit mehr Hyperparameter abstimmen.

In scikit-learn kannst du das mit RandomizedSearchCV machen. Es hat die gleiche API wie GridSearchCV, nur dass du statt konkreter Hyperparameterwerte eine Parameter-Verteilung angibst, aus der gesampelt wird. Probieren wir es aus! Die Parameterverteilung ist bereits für dich vorbereitet, ebenso ein Random-Forest-Klassifikator namens clf.

Diese Übung ist Teil des Kurses

<Kurs>Marketing Analytics: Kundenabwanderung in Python vorhersagen</Kurs>
Kurs ansehen

Interaktive praktische Übung

Versuche dich an dieser Übung, indem du diesen Beispielcode vervollständigst.

# Import RandomizedSearchCV
Code bearbeiten und ausführen