MCPcopy Create free account
hub / github.com/antmachineintelligence/mtgbmcode / TrainOneIter

Method TrainOneIter

src/boosting/rf.hpp:103–166  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

101 }
102
103 bool TrainOneIter(const score_t* gradients, const score_t* hessians) override {
104 // bagging logic
105 Bagging(iter_);
106 CHECK(gradients == nullptr);
107 CHECK(hessians == nullptr);
108
109 gradients = gradients_.data();
110 hessians = hessians_.data();
111 for (int cur_tree_id = 0; cur_tree_id < num_tree_per_iteration_; ++cur_tree_id) {
112 std::unique_ptr<Tree> new_tree(new Tree(2));
113 size_t offset = static_cast<size_t>(cur_tree_id)* num_data_;
114 if (class_need_train_[cur_tree_id]) {
115 auto grad = gradients + offset;
116 auto hess = hessians + offset;
117
118 // need to copy gradients for bagging subset.
119 if (is_use_subset_ && bag_data_cnt_ < num_data_) {
120 for (int i = 0; i < bag_data_cnt_; ++i) {
121 tmp_grad_[i] = grad[bag_data_indices_[i]];
122 tmp_hess_[i] = hess[bag_data_indices_[i]];
123 }
124 grad = tmp_grad_.data();
125 hess = tmp_hess_.data();
126 }
127
128 new_tree.reset(tree_learner_->Train(grad, hess, is_constant_hessian_,
129 forced_splits_json_));
130 }
131
132 if (new_tree->num_leaves() > 1) {
133 double pred = init_scores_[cur_tree_id];
134 auto residual_getter = [pred](const label_t* label, int i) {return static_cast<double>(label[i]) - pred; };
135 tree_learner_->RenewTreeOutput(new_tree.get(), objective_function_, residual_getter,
136 num_data_, bag_data_indices_.data(), bag_data_cnt_);
137 if (std::fabs(init_scores_[cur_tree_id]) > kEpsilon) {
138 new_tree->AddBias(init_scores_[cur_tree_id]);
139 }
140 // update score
141 MultiplyScore(cur_tree_id, (iter_ + num_init_iteration_));
142 UpdateScore(new_tree.get(), cur_tree_id);
143 MultiplyScore(cur_tree_id, 1.0 / (iter_ + num_init_iteration_ + 1));
144 } else {
145 // only add default score one-time
146 if (models_.size() < static_cast<size_t>(num_tree_per_iteration_)) {
147 double output = 0.0;
148 if (!class_need_train_[cur_tree_id]) {
149 if (objective_function_ != nullptr) {
150 output = objective_function_->BoostFromScore(cur_tree_id);
151 } else {
152 output = init_scores_[cur_tree_id];
153 }
154 }
155 new_tree->AsConstantTree(output);
156 MultiplyScore(cur_tree_id, (iter_ + num_init_iteration_));
157 UpdateScore(new_tree.get(), cur_tree_id);
158 MultiplyScore(cur_tree_id, 1.0 / (iter_ + num_init_iteration_ + 1));
159 }
160 }

Callers

nothing calls this directly

Calls 11

dataMethod · 0.80
resetMethod · 0.80
push_backMethod · 0.80
TrainMethod · 0.45
num_leavesMethod · 0.45
RenewTreeOutputMethod · 0.45
getMethod · 0.45
AddBiasMethod · 0.45
sizeMethod · 0.45
BoostFromScoreMethod · 0.45
AsConstantTreeMethod · 0.45

Tested by

no test coverage detected