ÎncepețiÎncepe gratuit

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

Vezi cursul

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ă predicted care 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 variabilei predicted.
  • Definește variabila responses pentru 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 în predicted.

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(___)
Editează și rulează codul