Standardizzare i dati
Alcuni modelli, come K-nearest neighbors (KNN) e le reti neurali, funzionano meglio con dati scalati, quindi standardizzeremo i nostri dati.
Rimuoveremo anche le variabili poco importanti (giorno della settimana), in base alle feature importances, indicizzando i DataFrame delle feature con .iloc[]. KNN usa le distanze per trovare punti simili per le previsioni, quindi le feature con valori più grandi pesano di più di quelle con valori piccoli. Scalare i dati risolve questo problema.
sklearn e la sua scale() standardizzano i dati, impostando la media a 0 e la deviazione standard a 1. Idealmente useremmo StandardScaler con fit_transform() sui dati di training e transform() sui dati di test, ma qui siamo limitati a 15 righe di codice.
Una volta scalati i dati, verificheremo che abbia funzionato tracciando gli istogrammi dei dati.
Questo esercizio fa parte del corso
Machine Learning per la finanza in Python
Istruzioni dell'esercizio
- Rimuovi le feature del giorno della settimana dalle feature di train/test usando
.iloc(i giorni della settimana sono le ultime 4 feature). - Standardizza
train_featuresetest_featuresusandoscale()di sklearn; salva le feature scalate comescaled_train_featuresescaled_test_features. - Traccia un istogramma della media mobile dell'RSI a 14 giorni (indicizzata a
[:, 2]) dalletrain_featuresnon scalate nel primo subplot (ax[0]). - Traccia un istogramma della media mobile dell'RSI a 14 giorni standardizzata nel secondo subplot (
ax[1]).
esercizio interattivo pratico
Prova questo esercizio completando questo codice di esempio.
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()