Стандартизация данных
Некоторые модели, например метод K ближайших соседей (KNN) и нейронные сети, лучше работают с масштабированными данными — поэтому мы стандартизируем наши данные.
Также мы удалим незначимые переменные (день недели) на основе важности признаков, отобрав нужные столбцы из DataFrames с помощью .iloc[]. KNN использует расстояния для поиска похожих точек при предсказании, поэтому признаки с большими значениями доминируют над малыми. Масштабирование данных устраняет эту проблему.
Функция scale() из sklearn выполняет стандартизацию данных: устанавливает среднее значение равным 0 и стандартное отклонение равным 1. В идеале следовало бы использовать StandardScaler с fit_transform() на обучающих данных и fit() на тестовых, однако здесь мы ограничены 15 строками кода.
После масштабирования данных проверьте результат, построив гистограммы.
Это упражнение является частью курса
Машинное обучение для финансов на Python
Инструкции к упражнению
- Удалите признаки дня недели из обучающей и тестовой выборок с помощью
.iloc(признаки дня недели — последние 4 столбца). - Стандартизируйте
train_featuresиtest_featuresс помощью функцииscale()из sklearn; сохраните результат в переменныхscaled_train_featuresиscaled_test_features. - Постройте гистограмму скользящего среднего RSI за 14 дней (индекс
[:, 2]) из нестандартизированныхtrain_featuresна первом подграфике (ax[0]). - Постройте гистограмму стандартизированного скользящего среднего RSI за 14 дней на втором подграфике (
ax[1]).
Интерактивное практическое упражнение
Попробуйте выполнить это упражнение, дополнив этот пример кода.
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()