Деревья с градиентным бустингом: предсказание
После того как модель обучена, следующий шаг — получить с её помощью предсказания. В отличие от базового R, где для этого используется функция predict(), в sparklyr применяется функция ml_predict(). Она принимает два аргумента: модель и тестовые данные.
ml_predict(a_model, testing_data)
Распространённая задача — сравнить предсказанные значения с фактическими и построить графики в R. Ниже приведён типичный шаблон кода для подготовки таких данных. Обратите внимание: добавление столбца с предсказаниями в настоящее время выполняется локально, поэтому сначала необходимо собрать результаты с помощью collect().
predicted_vs_actual <- testing_data %>%
select(actual) %>%
collect() %>%
mutate(predicted)
Это упражнение является частью курса
Введение в Spark с sparklyr на R
Инструкции к упражнению
Подключение к Spark уже создано и доступно как spark_conn. Тиблы, связанные с обучающим и тестовым наборами данных в Spark, предварительно определены как track_data_to_model_tbl и track_data_to_predict_tbl соответственно. Модель деревьев с градиентным бустингом предварительно определена как gradient_boosted_trees_model.
- Определите переменную
predicted, которая будет содержать предсказания модели для тестовых данных.- Вызовите
ml_predict(), передав модель и тестовые данные в качестве аргументов. Эта функция сгенерирует предсказания для тестового набора данных и добавит их в новый столбец с именемprediction. - С помощью
pull()извлеките этот столбец и присвойте его переменнойpredicted.
- Вызовите
- Определите переменную
responses, чтобы подготовить данные для сравнения предсказанных значений с фактическими:- Выберите столбец с ответами
year. - Соберите результаты с помощью
collect(). - Используйте
mutate(), чтобы добавить предсказания из переменнойpredicted.
- Выберите столбец с ответами
Интерактивное практическое упражнение
Попробуйте выполнить это упражнение, дополнив этот пример кода.
# Training, testing sets & model are pre-defined
track_data_to_model_tbl
track_data_to_predict_tbl
gradient_boosted_trees_model
# Predict the responses for the testing data
predicted <- ___(
___,
___) %>% pull(prediction)
# Prepare the data for comparing predicted responses with actual responses
responses <- track_data_to_predict_tbl %>%
# Select the response column
___ %>%
# Collect the results
___ %>%
# Add in the predictions
mutate(___)