Identifică adâncimea optimă a arborelui
Acum vei ajusta parametrul max_depth al arborelui de decizie pentru a descoperi valoarea care reduce supraadaptarea (overfitting), menținând în același timp performanțe bune ale modelului. Vei parcurge un ciclu for prin mai multe valori ale parametrului max_depth, vei antrena un arbore de decizie pentru fiecare valoare și vei calcula metricile de performanță.
Lista depth_list cu valorile candidate ale parametrului a fost deja încărcată. Tabloul depth_tuning a fost construit cu 2 coloane: prima conține valorile candidate pentru adâncime, iar a doua este un spațiu rezervat pentru scorul de recall. Variabilele de caracteristici și țintă au fost încărcate ca train_X, train_Y pentru datele de antrenament și test_X, test_Y pentru datele de test. Bibliotecile numpy și pandas sunt încărcate ca np și respectiv pd.
Acest exercițiu face parte din cursul
Machine Learning pentru Marketing în Python
Instrucțiuni pentru exercițiu
- Rulează un ciclu
forpeste intervalul de la 0 până la lungimea listeidepth_list. - Pentru fiecare valoare candidată a adâncimii, inițializează și antrenează un clasificator de tip arbore de decizie, apoi prezice abandonul pe datele de test.
- Pentru fiecare valoare candidată a adâncimii, calculează scorul de recall folosind funcția
recall_score()și stochează-l în a doua coloană a luidepth_tunning. - Creează un DataFrame
pandasdindepth_tuningcu numele de coloane corespunzătoare.
Exercițiu interactiv practic
Încearcă acest exercițiu completând acest cod de exemplu.
# Run a for loop over the range of depth list length
for index in ___(0, len(depth_list)):
# Initialize and fit decision tree with the `max_depth` candidate
mytree = DecisionTreeClassifier(___=depth_list[index])
mytree.fit(___, train_Y)
# Predict churn on the testing data
pred_test_Y = mytree.predict(___)
# Calculate the recall score
depth_tuning[index,1] = ___(test_Y, ___)
# Name the columns and print the array as pandas DataFrame
col_names = ['Max_Depth','Recall']
print(pd.DataFrame(depth_tuning, columns=___))