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
Instrucțiuni pentru exercițiu
- Creează o coloană
soybean_test$pred.lincu predicțiile din modelul liniarmodel.lin. - Creează o coloană
soybean_test$pred.gamcu predicțiile din modelul GAMmodel.gam.- Pentru modelele GAM, metoda
predict()returnează o matrice, așa că foloseșteas.numeric()pentru a converti matricea într-un vector.
- Pentru modelele GAM, metoda
- Completează spațiile libere pentru a aplica
pivot_longer()pe coloanele de predicții într-o singură coloană de valoripred, cu coloana cheiemodeltype. Numește noul cadru de date lungsoybean_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 deTime. - Grafice de tip punct-și-linie ale predicțiilor (
pred) în funcție deTime. - Observă că modelul liniar prezice uneori greutăți negative! Modelul GAM face același lucru?
- Un grafic de tip scatter al variabilei
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")