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
Instrucciones del ejercicio
- Crea una columna
soybean_test$pred.lincon las predicciones del modelo linealmodel.lin. - Crea una columna
soybean_test$pred.gamcon las predicciones del modelo GAMmodel.gam.- Para modelos GAM, el método
predict()devuelve una matriz, así que usaas.numeric()para convertirla a un vector.
- Para modelos GAM, el método
- Completa los espacios en blanco para hacer
pivot_longer()de las columnas de predicción en una única columna de valorespredcon una columna clavemodeltype. Llamasoybean_longal 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
weighten función deTime. - Gráficos de puntos y líneas de las predicciones (
pred) en función deTime. - ¡Fíjate en que el modelo lineal a veces predice pesos negativos! ¿Lo hace el modelo GAM?
- Un diagrama de dispersión de
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")