CommencezCommencez gratuitement

Prédire avec le modèle de soya sur les données de test

Dans cet exercice, vous allez appliquer les modèles de soya de l'exercice précédent (model.lin et model.gam, déjà chargés) à de nouvelles données : soybean_test.

Cette activité fait partie du cours

Apprentissage supervisé en R : régression

Voir le cours

Instructions de l’exercice

  • Créez une colonne soybean_test$pred.lin avec les prédictions du modèle linéaire model.lin.
  • Créez une colonne soybean_test$pred.gam avec les prédictions du modèle GAM model.gam.
    • Pour les modèles GAM, la méthode predict() retourne une matrice ; utilisez donc as.numeric() pour convertir la matrice en vecteur.
  • Remplissez les blancs pour utiliser pivot_longer() et convertir les colonnes de prédictions en une seule colonne de valeurs pred avec une colonne clé modeltype. Nommez la version longue du tableau de données soybean_long.
  • Calculez et comparez le RMSE des deux modèles.
    • Lequel s'en tire 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 weight en fonction de Time.
    • Des graphiques points-et-lignes des prédictions (pred) en fonction de Time.
    • Remarquez que le modèle linéaire prédit parfois des poids négatifs ! Le modèle GAM fait-il de même ?

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")
  
Modifier et exécuter le code