НачатьНачать бесплатно

Предсказание на тестовых данных с помощью модели для сои

В этом упражнении вы применяете модели для сои из предыдущего упражнения (model.lin и model.gam, уже загружены) к новым данным: soybean_test.

Это упражнение является частью курса

Обучение с учителем в R: регрессия

Посмотреть курс

Инструкции к упражнению

  • Создайте столбец soybean_test$pred.lin с предсказаниями линейной модели model.lin.
  • Создайте столбец soybean_test$pred.gam с предсказаниями GAM-модели model.gam.
    • Для GAM-моделей метод predict() возвращает матрицу, поэтому используйте as.numeric(), чтобы преобразовать матрицу в вектор.
  • Заполните пропуски, чтобы применить 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")
  
Редактировать и запускать код