Kom igångKom igång gratis

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

Visa kurs

Ö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 predicted som 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 namnet prediction.
    • Använd pull() för att extrahera kolumnen och tilldela den till predicted.
  • Definiera variabeln responses fö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ån predicted.

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(___)
Redigera och kör kod