НачатьНачать бесплатно

Деревья с градиентным бустингом: предсказание

После того как модель обучена, следующий шаг — получить с её помощью предсказания. В отличие от базового 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(___)
Редактировать и запускать код