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
Övningsinstruktioner
- Skapa en kolumn
soybean_test$pred.linmed förutsägelser från den linjära modellenmodel.lin. - Skapa en kolumn
soybean_test$pred.gammed förutsägelser från GAM-modellenmodel.gam.- För GAM-modeller returnerar
predict()en matris, så användas.numeric()för att konvertera matrisen till en vektor.
- För GAM-modeller returnerar
- Fyll i luckorna för att använda
pivot_longer()på förutsägelsekolumnerna till en enda värdekolumnpredmed nyckelkolumnenmodeltype. Kalla den långa dataramensoybean_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
weightsom funktion avTime. - Punkt-och-linjediagram av förutsägelserna (
pred) som funktion avTime. - Observera att den linjära modellen ibland förutsäger negativa vikter! Gör GAM-modellen det?
- Ett spridningsdiagram av
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")