ÎncepețiÎncepe gratuit

Generează predicții cu modelul pentru soia pe datele de test

În acest exercițiu, vei aplica modelele pentru soia din exercițiul anterior (model.lin și model.gam, deja încărcate) pe date noi: soybean_test.

Acest exercițiu face parte din cursul

Învățare supervizată în R: Regresia

Vezi cursul

Instrucțiuni pentru exercițiu

  • Creează o coloană soybean_test$pred.lin cu predicțiile din modelul liniar model.lin.
  • Creează o coloană soybean_test$pred.gam cu predicțiile din modelul GAM model.gam.
    • Pentru modelele GAM, metoda predict() returnează o matrice, așa că folosește as.numeric() pentru a converti matricea într-un vector.
  • Completează spațiile libere pentru a aplica pivot_longer() pe coloanele de predicții într-o singură coloană de valori pred, cu coloana cheie modeltype. Numește noul cadru de date lung soybean_long.
  • Calculează și compară RMSE-ul ambelor modele.
    • Care model are performanțe mai bune?
  • Rulează codul pentru a compara predicțiile fiecărui model cu valorile reale ale greutății medii a frunzelor.
    • Un grafic de tip scatter al variabilei weight în funcție de Time.
    • Grafice de tip punct-și-linie ale predicțiilor (pred) în funcție de Time.
    • Observă că modelul liniar prezice uneori greutăți negative! Modelul GAM face același lucru?

Exercițiu interactiv practic

Încearcă acest exercițiu completând acest cod de exemplu.

# 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")
  
Editează și rulează codul