Arbres à gradient boosting : prédiction
Une fois votre modèle entraîné, l'étape suivante consiste à produire des prédictions. Contrairement à base R, qui utilise la fonction predict() pour prédire, sparklyr utilise la fonction ml_predict(). ml_predict() prend deux arguments : un modèle et des données de test.
ml_predict(a_model, testing_data)
Un cas d'usage courant consiste à comparer les réponses prédites aux réponses réelles, puis à tracer ces valeurs dans R. Le canevas de code pour préparer ces données est le suivant. Notez qu'à l'heure actuelle, l'ajout d'une colonne de prédiction doit se faire localement ; vous devez donc d'abord récupérer les résultats avec collect().
predicted_vs_actual <- testing_data %>%
select(actual) %>%
collect() %>%
mutate(predicted)
Cette activité fait partie du cours
Introduction à Spark avec sparklyr en R
Instructions de l’exercice
Une connexion Spark a été créée pour vous sous le nom spark_conn. Les tibbles associés aux ensembles de données d'entraînement et de test stockés dans Spark ont été prédéfinis comme track_data_to_model_tbl et track_data_to_predict_tbl, respectivement. Le modèle d'arbres à gradient boosting a été prédéfini comme gradient_boosted_trees_model.
- Définissez une variable
predictedqui contient les prédictions du modèle pour nos données de test.- Appelez
ml_predict()avec le modèle et les données de test en arguments. Cette fonction générera des prédictions pour l'ensemble de test et les ajoutera comme nouvelle colonne nomméeprediction. - À l'aide de
pull(), vous pouvez extraire cette colonne et l'assigner àpredicted.
- Appelez
- Définissez la variable
responsespour préparer les données en vue de comparer les réponses prédites aux réponses réelles :- Sélectionnez la colonne de réponse
year. - Récupérez les résultats avec
collect(). - Utilisez
mutate()pour ajouter les prédictions contenues danspredicted.
- Sélectionnez la colonne de réponse
Exercice interactif pratique
Essayez cet exercice en complétant ce code d’exemple.
# 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(___)