Предсказание на тестовых данных с помощью модели для сои
В этом упражнении вы применяете модели для сои из предыдущего упражнения (model.lin и model.gam, уже загружены) к новым данным: soybean_test.
Это упражнение является частью курса
Обучение с учителем в R: регрессия
Инструкции к упражнению
- Создайте столбец
soybean_test$pred.linс предсказаниями линейной моделиmodel.lin. - Создайте столбец
soybean_test$pred.gamс предсказаниями GAM-моделиmodel.gam.- Для GAM-моделей метод
predict()возвращает матрицу, поэтому используйтеas.numeric(), чтобы преобразовать матрицу в вектор.
- Для GAM-моделей метод
- Заполните пропуски, чтобы применить
pivot_longer()к столбцам с предсказаниями: объедините их в один столбец значенийpredс ключевым столбцомmodeltype. Назовите результирующий длинный фрейм данныхsoybean_long. - Вычислите RMSE для обеих моделей и сравните результаты.
- Какая модель показывает лучший результат?
- Запустите код, чтобы сравнить предсказания каждой модели с фактическими средними значениями веса листьев.
- Диаграмма рассеяния
weightкак функцииTime. - Точечно-линейные графики предсказаний (
pred) как функцииTime. - Обратите внимание, что линейная модель иногда предсказывает отрицательные значения веса! Делает ли то же самое GAM-модель?
- Диаграмма рассеяния
Интерактивное практическое упражнение
Попробуйте выполнить это упражнение, дополнив этот пример кода.
# 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")