CommencerCommencez gratuitement

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>
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() renvoie une matrice ; utilisez donc as.numeric() pour convertir la matrice en vecteur.
  • Complétez les blancs pour utiliser pivot_longer() et rassembler les colonnes de prédiction en une seule colonne de valeurs pred avec une colonne clé modeltype. Appelez le DataFrame long soybean_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 weight en fonction de Time.
    • Des graphes 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 le fait-il ?

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