Losowe przeszukiwanie siatki
# Call GridSearchCV
grid_search = GridSearchCV(clf, param_grid)
# Fit the model
grid_search.fit(X, y)
W powyższym fragmencie kodu z poprzedniego ćwiczenia mogło się rzucić w oczy, że pierwsza linia wykonała się niemal natychmiast, podczas gdy wywołanie .fit() zajęło kilka sekund.
To dlatego, że to właśnie .fit() przeprowadza właściwe przeszukiwanie siatki – w naszym przypadku obejmowało ono wiele różnych kombinacji. Im większa siatka hiperparametrów, tym wolniejsze przeszukiwanie. Żeby rozwiązać ten problem, zamiast sprawdzać każdą możliwą kombinację, można losowo poruszać się po siatce i testować wybrane zestawy wartości. Istnieje niewielkie ryzyko, że pominiemy najlepszą kombinację, ale za to zaoszczędzimy dużo czasu – lub zdołamy dostroić więcej hiperparametrów w tym samym czasie.
W scikit-learn służy do tego RandomizedSearchCV. Ma takie samo API jak GridSearchCV, z tą różnicą, że zamiast konkretnych wartości hiperparametrów podajesz rozkład parametrów, z którego algorytm może losować próbki. Przetestuj to teraz! Rozkład parametrów oraz klasyfikator lasu losowego clf zostały już przygotowane.
To ćwiczenie jest częścią kursu
Marketing Analytics: Przewidywanie rezygnacji klientów w Pythonie
Interaktywne ćwiczenie praktyczne
Spróbuj tego ćwiczenia, uzupełniając ten przykładowy kod.
# Import RandomizedSearchCV