Mit dem Soybean-Modell auf Testdaten vorhersagen
In dieser Übung wendest du die Soybean-Modelle aus der vorherigen Aufgabe (model.lin und model.gam, bereits geladen) auf neue Daten an: soybean_test.
Diese Übung ist Teil des Kurses
<Kurs>Überwachtes Lernen in R: Regression</Kurs>Übungsanweisungen
- Erstelle eine Spalte
soybean_test$pred.linmit Vorhersagen aus dem linearen Modellmodel.lin. - Erstelle eine Spalte
soybean_test$pred.gammit Vorhersagen aus dem GAM-Modellmodel.gam.- Für GAM-Modelle gibt die
predict()-Methode eine Matrix zurück, daher mitas.numeric()die Matrix in einen Vektor umwandeln.
- Für GAM-Modelle gibt die
- Ergänze die Lücken, um mit
pivot_longer()die Vorhersagespalten in eine einzelne Wert-Spaltepredmit der Schlüsselspaltemodeltypezu überführen. Nenne den langen Data Framesoybean_long. - Berechne und vergleiche die RMSE beider Modelle.
- Welches Modell schneidet besser ab?
- Führe den Code aus, um die Vorhersagen jedes Modells mit den tatsächlichen durchschnittlichen Blattgewichten zu vergleichen.
- Ein Streudiagramm von
weightin Abhängigkeit vonTime. - Punkt-und-Linien-Diagramme der Vorhersagen (
pred) in Abhängigkeit vonTime. - Beachte: Das lineare Modell sagt manchmal negative Gewichte vorher! Tut das GAM-Modell das auch?
- Ein Streudiagramm von
Interaktive praktische Übung
Versuche dich an dieser Übung, indem du diesen Beispielcode vervollständigst.
# 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")