Kom igångKom igång gratis

Förutsäg med sojabönmodellen på testdata

I den här övningen ska du tillämpa sojabönmodellerna från föregående övning (model.lin och model.gam, redan inlästa) på nya data: soybean_test.

Den här övningen är en del av kursen

Övervakad inlärning i R: Regression

Visa kurs

Övningsinstruktioner

  • Skapa en kolumn soybean_test$pred.lin med förutsägelser från den linjära modellen model.lin.
  • Skapa en kolumn soybean_test$pred.gam med förutsägelser från GAM-modellen model.gam.
    • För GAM-modeller returnerar predict() en matris, så använd as.numeric() för att konvertera matrisen till en vektor.
  • Fyll i luckorna för att använda pivot_longer() på förutsägelsekolumnerna till en enda värdekolumn pred med nyckelkolumnen modeltype. Kalla den långa dataramen soybean_long.
  • Beräkna och jämför RMSE för båda modellerna.
    • Vilken modell presterar bäst?
  • Kör koden för att jämföra varje modells förutsägelser mot de faktiska genomsnittliga bladavikterna.
    • Ett spridningsdiagram av weight som funktion av Time.
    • Punkt-och-linjediagram av förutsägelserna (pred) som funktion av Time.
    • Observera att den linjära modellen ibland förutsäger negativa vikter! Gör GAM-modellen det?

Interaktiv övning med praktiskt arbete

Testa den här övningen genom att slutföra den här exempelkoden.

# soybean_test is available
summary(soybean_test)

# Get predictions from linear model
soybean_test$pred.lin <- ___(___, newdata = ___)

# Get predictions from gam model
soybean_test$pred.gam <- ___(___(___, newdata = ___))

# Pivot the predictions into a "long" dataset
soybean_long <- soybean_test %>%
  pivot_longer(cols = c(___, ___), names_to = ___, values_to = ___)

# Calculate the rmse
soybean_long %>%
  mutate(residual = weight - pred) %>%     # residuals
  group_by(modeltype) %>%                  # group by modeltype
  summarize(rmse = ___(___(___))) # calculate the RMSE

# Compare the predictions against actual weights on the test data
soybean_long %>%
  ggplot(aes(x = Time)) +                          # the column for the x axis
  geom_point(aes(y = weight)) +                    # the y-column for the scatterplot
  geom_point(aes(y = pred, color = modeltype)) +   # the y-column for the point-and-line plot
  geom_line(aes(y = pred, color = modeltype, linetype = modeltype)) + # the y-column for the point-and-line plot
  scale_color_brewer(palette = "Dark2")
  
Redigera och kör kod