Arbori cu gradient boosting: predicție
După ce ai rulat modelul, următorul pas este să faci predicții cu el. Spre deosebire de R de bază, care folosește funcția predict() pentru predicții, sparklyr folosește funcția ml_predict(). Aceasta primește două argumente: un model și datele de testare.
ml_predict(a_model, testing_data)
Un scenariu frecvent este să compari răspunsurile prezise cu cele reale și să le vizualizezi în R. Tiparul de cod pentru pregătirea acestor date este următorul. Reține că, în prezent, adăugarea unei coloane de predicție trebuie făcută local, deci trebuie să colectezi mai întâi rezultatele.
predicted_vs_actual <- testing_data %>%
select(actual) %>%
collect() %>%
mutate(predicted)
Acest exercițiu face parte din cursul
Introducere în Spark cu sparklyr în R
Instrucțiuni pentru exercițiu
O conexiune Spark a fost creată pentru tine sub numele spark_conn. Tibble-urile atașate seturilor de date de antrenament și de testare stocate în Spark au fost predefinite ca track_data_to_model_tbl, respectiv track_data_to_predict_tbl. Modelul cu arbori cu gradient boosting a fost predefinit ca gradient_boosted_trees_model.
- Definește o variabilă
predictedcare să conțină predicțiile modelului pentru datele de testare.- Apelează
ml_predict()cu modelul și datele de testare ca argumente. Această funcție va genera predicții pentru setul de date de testare și le va adăuga ca o nouă coloană numităprediction. - Folosind
pull(), poți extrage această coloană și o poți atribui variabileipredicted.
- Apelează
- Definește variabila
responsespentru a pregăti datele în vederea comparării răspunsurilor prezise cu cele reale:- Selectează coloana cu răspunsul,
year. - Colectează rezultatele.
- Folosește
mutate()pentru a adăuga predicțiile stocate înpredicted.
- Selectează coloana cu răspunsul,
Exercițiu interactiv practic
Încearcă acest exercițiu completând acest cod de exemplu.
# 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(___)