Prédire avec le modèle soybean sur des données de test
Dans cet exercice, vous allez appliquer les modèles soybean de l’exercice précédent (model.lin et model.gam, déjà chargés) à de nouvelles données : soybean_test.
Cet exercice fait partie du cours
<cours>Apprentissage supervisé en R : Régression</cours>Instructions de l’exercice
- Créez une colonne
soybean_test$pred.linavec les prédictions du modèle linéairemodel.lin. - Créez une colonne
soybean_test$pred.gamavec les prédictions du modèle GAMmodel.gam.- Pour les modèles GAM, la méthode
predict()renvoie une matrice ; utilisez doncas.numeric()pour convertir la matrice en vecteur.
- Pour les modèles GAM, la méthode
- Complétez les blancs pour utiliser
pivot_longer()et rassembler les colonnes de prédiction en une seule colonne de valeurspredavec une colonne clémodeltype. Appelez le DataFrame longsoybean_long. - Calculez et comparez le RMSE des deux modèles.
- Lequel s’en sort le mieux ?
- Exécutez le code pour comparer les prédictions de chaque modèle aux poids moyens réels des feuilles.
- Un nuage de points de
weighten fonction deTime. - Des graphes points-et-lignes des prédictions (
pred) en fonction deTime. - Remarquez que le modèle linéaire prédit parfois des poids négatifs ! Le modèle GAM le fait-il ?
- Un nuage de points de
Exercice interactif pratique
Essayez cet exercice en complétant ce code d’exemple.
# 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")