Gradient boosted trees: dự đoán
Khi đã chạy xong mô hình, bước tiếp theo là dùng nó để tạo dự đoán. Khác với base R vốn dùng hàm predict() để dự đoán, sparklyr dùng hàm ml_predict(). ml_predict() nhận hai đối số: một mô hình và một bộ dữ liệu kiểm thử.
ml_predict(a_model, testing_data)
Một trường hợp sử dụng phổ biến là so sánh giá trị dự đoán với giá trị thực tế, sau đó bạn có thể vẽ biểu đồ trong R. Mẫu mã để chuẩn bị dữ liệu như sau. Lưu ý hiện tại việc thêm một cột dự đoán phải thực hiện cục bộ, nên bạn cần collect() kết quả trước.
predicted_vs_actual <- testing_data %>%
select(actual) %>%
collect() %>%
mutate(predicted)
Bài tập này là một phần của khóa học
Nhập môn Spark với sparklyr trong R
Hướng dẫn bài tập
Kết nối Spark đã được tạo sẵn là spark_conn. Các tibble gắn với tập huấn luyện và tập kiểm thử lưu trên Spark đã được định nghĩa sẵn lần lượt là track_data_to_model_tbl và track_data_to_predict_tbl. Mô hình gradient boosted trees đã được định nghĩa sẵn là gradient_boosted_trees_model.
- Định nghĩa biến
predictedchứa các dự đoán của mô hình cho dữ liệu kiểm thử.- Gọi
ml_predict()với mô hình và dữ liệu kiểm thử làm đối số. Hàm này sẽ tạo dự đoán cho tập kiểm thử và thêm chúng thành một cột mới tên làprediction. - Dùng
pull()để trích xuất cột này và gán chopredicted.
- Gọi
- Định nghĩa biến
responsesđể chuẩn bị dữ liệu so sánh giữa dự đoán và giá trị thực tế:- Chọn cột phản hồi
year. - Thu thập kết quả (
collect). - Dùng
mutate()để thêm các dự đoán trongpredicted.
- Chọn cột phản hồi
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.
# Training, testing sets & model are pre-defined
track_data_to_model_tbl
track_data_to_predict_tbl
gradient_boosted_trees_model
# Predict the responses for the testing data
predicted <- ___(
___,
___) %>% pull(prediction)
# Prepare the data for comparing predicted responses with actual responses
responses <- track_data_to_predict_tbl %>%
# Select the response column
___ %>%
# Collect the results
___ %>%
# Add in the predictions
mutate(___)