Ein Modell mit Test-/Train-Split evaluieren
Jetzt testest du das Modell mpg_model auf den Testdaten mpg_test.
Die Funktionen rmse() und r_squared() zum Berechnen von RMSE und R-squared wurden der Einfachheit halber bereitgestellt:
rmse(predcol, ycol)
r_squared(predcol, ycol)
wobei:
- predcol: die vorhergesagten Werte
- ycol: das tatsächliche Ergebnis
Außerdem wirst du die Vorhersagen gegen das tatsächliche Ergebnis plotten.
Im Allgemeinen ist die Modellleistung auf den Trainingsdaten besser als auf den Testdaten (manchmal hat das Test-Set allerdings „Glück“). Ein kleiner Leistungsunterschied ist in Ordnung; wenn die Leistung auf dem Training jedoch deutlich besser ist, gibt es ein Problem.
Die Data Frames mpg_train und mpg_test sowie das Modell mpg_model sind zusammen mit den Funktionen rmse() und r_squared() bereits geladen.
Diese Übung ist Teil des Kurses
<Kurs>Überwachtes Lernen in R: Regression</Kurs>Übungsanweisungen
- Sage die städtische Kraftstoffeffizienz aus
hwyauf denmpg_train-Daten voraus. Weise die Vorhersagen der Spaltepredzu. - Sage die städtische Kraftstoffeffizienz aus
hwyauf denmpg_test-Daten voraus. Weise die Vorhersagen der Spaltepredzu. - Verwende
rmse(), um das RMSE für Test- und Trainingsdaten zu berechnen. Vergleiche: Sind die Leistungen ähnlich? - Mache dasselbe mit
r_squared(). Sind die Leistungen ähnlich? - Verwende
ggplot2, um die Vorhersagen gegenctyauf dentest-Daten zu plotten.
Interaktive praktische Übung
Versuche dich an dieser Übung, indem du diesen Beispielcode vervollständigst.
# Examine the objects that have been loaded
ls.str()
# predict cty from hwy for the training set
mpg_train$pred <- ___
# predict cty from hwy for the test set
mpg_test$pred <- ___
# Evaluate the rmse on both training and test data and print them
(rmse_train <- ___)
(rmse_test <- ___)
# Evaluate the r-squared on both training and test data.and print them
(rsq_train <- ___)
(rsq_test <- ___)
# Plot the predictions (on the x-axis) against the outcome (cty) on the test data
ggplot(___, aes(x = ___, y = ___)) +
geom_point() +
geom_abline()