Predykcja na danych testowych z użyciem modelu sojowego
W tym ćwiczeniu zastosujesz modele sojowe z poprzedniego ćwiczenia (model.lin i model.gam, już wczytane) do nowych danych: soybean_test.
To ćwiczenie jest częścią kursu
Nadzorowane uczenie maszynowe w R: regresja
Instrukcje do ćwiczenia
- Utwórz kolumnę
soybean_test$pred.linz predykcjami modelu liniowegomodel.lin. - Utwórz kolumnę
soybean_test$pred.gamz predykcjami modelu GAMmodel.gam.- W przypadku modeli GAM metoda
predict()zwraca macierz – użyj funkcjias.numeric(), aby przekonwertować ją na wektor.
- W przypadku modeli GAM metoda
- Uzupełnij puste miejsca, aby za pomocą
pivot_longer()przekształcić kolumny z predykcjami w jedną kolumnę wartościpredz kolumną kluczymodeltype. Nazwij nową ramkę danychsoybean_long. - Oblicz i porównaj RMSE obu modeli.
- Który model radzi sobie lepiej?
- Uruchom kod, aby porównać predykcje każdego modelu z rzeczywistymi średnimi wagami liści.
- Wykres punktowy wagi
weightjako funkcji czasuTime. - Wykresy punktowo-liniowe predykcji (
pred) jako funkcji czasuTime. - Zwróć uwagę, że model liniowy czasem przewiduje ujemne wagi! Czy model GAM też tak robi?
- Wykres punktowy wagi
Interaktywne ćwiczenie praktyczne
Spróbuj tego ćwiczenia, uzupełniając ten przykładowy kod.
# 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")