| 166 | } |
| 167 | |
| 168 | void RollbackOneIter() override { |
| 169 | if (iter_ <= 0) { return; } |
| 170 | int cur_iter = iter_ + num_init_iteration_ - 1; |
| 171 | // reset score |
| 172 | for (int cur_tree_id = 0; cur_tree_id < num_tree_per_iteration_; ++cur_tree_id) { |
| 173 | auto curr_tree = cur_iter * num_tree_per_iteration_ + cur_tree_id; |
| 174 | models_[curr_tree]->Shrinkage(-1.0); |
| 175 | MultiplyScore(cur_tree_id, (iter_ + num_init_iteration_)); |
| 176 | train_score_updater_->AddScore(models_[curr_tree].get(), cur_tree_id); |
| 177 | for (auto& score_updater : valid_score_updater_) { |
| 178 | score_updater->AddScore(models_[curr_tree].get(), cur_tree_id); |
| 179 | } |
| 180 | MultiplyScore(cur_tree_id, 1.0f / (iter_ + num_init_iteration_ - 1)); |
| 181 | } |
| 182 | // remove model |
| 183 | for (int cur_tree_id = 0; cur_tree_id < num_tree_per_iteration_; ++cur_tree_id) { |
| 184 | models_.pop_back(); |
| 185 | } |
| 186 | --iter_; |
| 187 | } |
| 188 | |
| 189 | void MultiplyScore(const int cur_tree_id, double val) { |
| 190 | train_score_updater_->MultiplyScore(val, cur_tree_id); |