| 721 | } |
| 722 | |
| 723 | void GBDT::GetPredictAt(int data_idx, double* out_result, int64_t* out_len) { |
| 724 | CHECK(data_idx >= 0 && data_idx <= static_cast<int>(valid_score_updater_.size())); |
| 725 | |
| 726 | const double* raw_scores = nullptr; |
| 727 | data_size_t num_data = 0; |
| 728 | if (data_idx == 0) { |
| 729 | raw_scores = GetTrainingScore(out_len); |
| 730 | num_data = train_score_updater_->num_data(); |
| 731 | |
| 732 | } else { |
| 733 | auto used_idx = data_idx - 1; |
| 734 | raw_scores = valid_score_updater_[used_idx]->score(); |
| 735 | num_data = valid_score_updater_[used_idx]->num_data(); |
| 736 | *out_len = static_cast<int64_t>(num_data) * num_class_ * num_labels_; |
| 737 | } |
| 738 | if (objective_function_ != nullptr) { |
| 739 | #pragma omp parallel for schedule(static) |
| 740 | for (data_size_t i = 0; i < num_data; ++i) { |
| 741 | std::vector<double> tree_pred(num_tree_per_iteration_); |
| 742 | for (int j = 0; j < num_tree_per_iteration_; ++j) { |
| 743 | tree_pred[j] = raw_scores[j * num_data + i]; |
| 744 | } |
| 745 | std::vector<double> tmp_result(num_class_); |
| 746 | objective_function_->ConvertOutput(tree_pred.data(), tmp_result.data()); |
| 747 | for (int j = 0; j < num_class_; ++j) { |
| 748 | out_result[j * num_data + i] = static_cast<double>(tmp_result[j]); |
| 749 | } |
| 750 | } |
| 751 | } else { |
| 752 | #pragma omp parallel for schedule(static) |
| 753 | for (data_size_t i = 0; i < num_data; ++i) { |
| 754 | for (int j = 0; j < num_tree_per_iteration_; ++j) { |
| 755 | for (int k = 0; k < num_labels_; ++k) { |
| 756 | out_result[k * num_tree_per_iteration_ * num_data + j * num_data + i] = static_cast<double>(raw_scores[k * num_tree_per_iteration_ * num_data + j * num_data + i]); |
| 757 | } |
| 758 | } |
| 759 | } |
| 760 | } |
| 761 | } |
| 762 | |
| 763 | void GBDT::ResetTrainingData(const Dataset* train_data, const ObjectiveFunction* objective_function, |
| 764 | const std::vector<const Metric*>& training_metrics) { |
nothing calls this directly
no test coverage detected