Zacznij terazZacznij za darmo

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

Zobacz kurs

Instrukcje do ćwiczenia

  • Utwórz kolumnę soybean_test$pred.lin z predykcjami modelu liniowego model.lin.
  • Utwórz kolumnę soybean_test$pred.gam z predykcjami modelu GAM model.gam.
    • W przypadku modeli GAM metoda predict() zwraca macierz – użyj funkcji as.numeric(), aby przekonwertować ją na wektor.
  • Uzupełnij puste miejsca, aby za pomocą pivot_longer() przekształcić kolumny z predykcjami w jedną kolumnę wartości pred z kolumną kluczy modeltype. Nazwij nową ramkę danych soybean_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 weight jako funkcji czasu Time.
    • Wykresy punktowo-liniowe predykcji (pred) jako funkcji czasu Time.
    • Zwróć uwagę, że model liniowy czasem przewiduje ujemne wagi! Czy model GAM też tak robi?

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")
  
Edytuj i uruchom kod