Data standaardiseren
Sommige modellen, zoals K-nearest neighbors (KNN) en neural networks, werken beter met geschaalde data — daarom standaardiseren we onze data.
We verwijderen ook onbelangrijke variabelen (dag van de week) op basis van feature-importance door de features-DataFrames te indexeren met .iloc[]. KNN gebruikt afstanden om vergelijkbare punten voor voorspellingen te vinden, dus grote features wegen zwaarder dan kleine. Door te schalen los je dat op.
skslearn's scale() zal data standaardiseren, waarbij het gemiddelde op 0 en de standaardafwijking op 1 wordt gezet. Idealiter gebruiken we StandardScaler met fit_transform() op de trainingsdata en fit() op de testdata, maar we zijn hier beperkt tot 15 regels code.
Zodra we de data hebben geschaald, controleren we of het gewerkt heeft door histogrammen van de data te plotten.
Deze oefening maakt deel uit van de cursus
Machine Learning voor finance in Python
Oefeninstructies
- Verwijder de dag-van-de-week-features uit de train/test-features met
.iloc(dag van de week zijn de laatste 4 features). - Standaardiseer
train_featuresentest_featuresmet sklearnsscale(); sla de geschaalde features op alsscaled_train_featuresenscaled_test_features. - Plot een histogram van het 14-daags RSI-voortschrijdend gemiddelde (geïndexeerd op
[:, 2]) uit ongeschaaldetrain_featuresop de eerste subplot (ax[0]]). - Plot een histogram van het gestandaardiseerde 14-daags RSI-voortschrijdend gemiddelde op de tweede subplot (
ax[1]).
Interactieve oefening met praktijkervaring
Probeer deze oefening door deze voorbeeldcode aan te vullen.
from sklearn.preprocessing import scale
# Remove unimportant features (weekdays)
train_features = train_features.iloc[:, :-4]
test_features = test_features.____
# Standardize the train and test features
scaled_train_features = scale(train_features)
scaled_test_features = ____
# Plot histograms of the 14-day SMA RSI before and after scaling
f, ax = plt.subplots(nrows=2, ncols=1)
train_features.iloc[:, 2].hist(ax=____)
ax[1].hist(scaled_train_features[:, 2])
plt.show()