EmpezarEmpieza gratis

Predecir con el modelo de soja en datos de prueba

En este ejercicio, aplicarás los modelos de soja del ejercicio anterior (model.lin y model.gam, ya cargados) a nuevos datos: soybean_test.

Este ejercicio forma parte del curso

Aprendizaje supervisado en R: Regresión

Ver curso

Instrucciones del ejercicio

  • Crea una columna soybean_test$pred.lin con las predicciones del modelo lineal model.lin.
  • Crea una columna soybean_test$pred.gam con las predicciones del modelo GAM model.gam.
    • Para modelos GAM, el método predict() devuelve una matriz, así que usa as.numeric() para convertirla a un vector.
  • Completa los espacios en blanco para hacer pivot_longer() de las columnas de predicción en una única columna de valores pred con una columna clave modeltype. Llama soybean_long al data frame en formato largo.
  • Calcula y compara el RMSE de ambos modelos.
    • ¿Qué modelo lo hace mejor?
  • Ejecuta el código para comparar las predicciones de cada modelo con los pesos medios de hoja reales.
    • Un diagrama de dispersión de weight en función de Time.
    • Gráficos de puntos y líneas de las predicciones (pred) en función de Time.
    • ¡Fíjate en que el modelo lineal a veces predice pesos negativos! ¿Lo hace el modelo GAM?

ejercicio interactivo práctico

Prueba este ejercicio completando este código de ejemplo.

# 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")
  
Editar y ejecutar código