Bắt đầu ngayBắt đầu miễn phí

Dự đoán với mô hình soybean trên dữ liệu kiểm tra

Trong bài tập này, bạn sẽ áp dụng các mô hình soybean từ bài trước (model.linmodel.gam, đã được nạp) lên dữ liệu mới: soybean_test.

Bài tập này là một phần của khóa học

Học có giám sát với R: Hồi quy

Xem khóa học

Hướng dẫn bài tập

  • Tạo cột soybean_test$pred.lin với dự đoán từ mô hình tuyến tính model.lin.
  • Tạo cột soybean_test$pred.gam với dự đoán từ mô hình gam model.gam.
    • Với mô hình GAM, phương thức predict() trả về một ma trận, nên dùng as.numeric() để chuyển ma trận thành vector.
  • Điền vào chỗ trống để pivot_longer() các cột dự đoán thành một cột giá trị duy nhất pred với cột khóa modeltype. Gọi khung dữ liệu dạng dài là soybean_long.
  • Tính và so sánh RMSE của cả hai mô hình.
    • Mô hình nào làm tốt hơn?
  • Chạy mã để so sánh dự đoán của mỗi mô hình với trọng lượng lá trung bình thực tế.
    • Biểu đồ scatter của weight theo Time.
    • Biểu đồ điểm-và-đường của các dự đoán (pred) theo Time.
    • Lưu ý rằng mô hình tuyến tính đôi khi dự đoán trọng lượng âm! Còn mô hình gam thì sao?

Bài tập tương tác thực hành trực tiếp

Hãy thử làm bài tập này bằng cách hoàn thành đoạn mã mẫu này.

# 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")
  
Chỉnh sửa và Chạy Mã