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.lin và model.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
Hướng dẫn bài tập
- Tạo cột
soybean_test$pred.linvới dự đoán từ mô hình tuyến tínhmodel.lin. - Tạo cột
soybean_test$pred.gamvới dự đoán từ mô hình gammodel.gam.- Với mô hình GAM, phương thức
predict()trả về một ma trận, nên dùngas.numeric()để chuyển ma trận thành vector.
- Với mô hình GAM, phương thức
- Đ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ấtpredvới cột khóamodeltype. 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
weighttheoTime. - Biểu đồ điểm-và-đường của các dự đoán (
pred) theoTime. - 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?
- Biểu đồ scatter của
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")