Gradient boosted trees: predikce
Jakmile model natrénuješ, dalším krokem je provést predikci. Na rozdíl od základního R, kde se k předpovědím používá funkce predict(), sparklyr pracuje s funkcí ml_predict(). Ta přijímá dva argumenty: model a testovací data.
ml_predict(a_model, testing_data)
Častým postupem je porovnání předpovězených hodnot se skutečnými, které pak můžeš vizualizovat přímo v R. Vzor kódu pro přípravu takových dat vypadá takto. Všimni si, že přidání sloupce s predikcí musí v současnosti proběhnout lokálně, takže je nejprve nutné výsledky shromáždit pomocí collect().
predicted_vs_actual <- testing_data %>%
select(actual) %>%
collect() %>%
mutate(predicted)
Toto cvičení je součástí kurzu
Úvod do Sparku se sparklyr v R
Pokyny k cvičení
Spark připojení je k dispozici jako spark_conn. Tibbles napojené na trénovací a testovací datové sady uložené ve Sparku jsou předdefinovány jako track_data_to_model_tbl a track_data_to_predict_tbl. Model gradient boosted trees je předdefinován jako gradient_boosted_trees_model.
- Definuj proměnnou
predicted, která bude obsahovat predikce modelu pro testovací data.- Zavolej
ml_predict()s modelem a testovacími daty jako argumenty. Tato funkce vygeneruje predikce pro testovací datovou sadu a přidá je jako nový sloupec s názvemprediction. - Pomocí
pull()tento sloupec extrahuj a přiřaď ho dopredicted.
- Zavolej
- Definuj proměnnou
responsespro přípravu dat k porovnání předpovězených a skutečných hodnot:- Vyber sloupec s odpovědí
year. - Shromáždi výsledky pomocí
collect(). - Pomocí
mutate()přidej predikce uložené vpredicted.
- Vyber sloupec s odpovědí
Interaktivní cvičení na vyzkoušení si v praxi
Vyzkoušejte si toto cvičení dokončením tohoto ukázkového kódu.
# 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(___)