Gradientförstärkta träd: prediktion
När du har kört din modell är nästa steg att göra en prediktion med den. Till skillnad från bas-R, som använder funktionen predict(), använder sparklyr funktionen ml_predict(). ml_predict() tar två argument: en modell och testdata.
ml_predict(a_model, testing_data)
Ett vanligt användningsfall är att jämföra de predikterade värdena med de faktiska värdena, vilket du kan visualisera med diagram i R. Koddmönstret för att förbereda dessa data ser ut så här. Observera att det för tillfället krävs att du lägger till en prediktionskolumn lokalt, så du måste samla in resultaten först.
predicted_vs_actual <- testing_data %>%
select(actual) %>%
collect() %>%
mutate(predicted)
Den här övningen är en del av kursen
Introduktion till Spark med sparklyr i R
Övningsinstruktioner
En Spark-anslutning har skapats åt dig som spark_conn. Tibbles kopplade till tränings- och testdatamängderna i Spark har fördefinierat som track_data_to_model_tbl respektive track_data_to_predict_tbl. Modellen för gradientförstärkta träd har fördefinierat som gradient_boosted_trees_model.
- Definiera en variabel
predictedsom innehåller modellens prediktioner för testdata.- Anropa
ml_predict()med modellen och testdata som argument. Den här funktionen genererar prediktioner för testdatamängden och lägger till dem som en ny kolumn med namnetprediction. - Använd
pull()för att extrahera kolumnen och tilldela den tillpredicted.
- Anropa
- Definiera variabeln
responsesför att förbereda data inför jämförelsen av predikterade och faktiska värden:- Välj responsvariabelkolumnen
year. - Samla in resultaten.
- Använd
mutate()för att lägga till prediktion frånpredicted.
- Välj responsvariabelkolumnen
Interaktiv övning med praktiskt arbete
Testa den här övningen genom att slutföra den här exempelkoden.
# 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(___)