| 53 | } |
| 54 | |
| 55 | int32 BoostedTreesEnsembleResource::next_node( |
| 56 | const int32 tree_id, const int32 node_id, const int32 index_in_batch, |
| 57 | const std::vector<TTypes<int32>::ConstVec>& bucketized_features) const { |
| 58 | DCHECK_LT(tree_id, tree_ensemble_->trees_size()); |
| 59 | DCHECK_LT(node_id, tree_ensemble_->trees(tree_id).nodes_size()); |
| 60 | const auto& node = tree_ensemble_->trees(tree_id).nodes(node_id); |
| 61 | |
| 62 | switch (node.node_case()) { |
| 63 | case boosted_trees::Node::kBucketizedSplit: { |
| 64 | const auto& split = node.bucketized_split(); |
| 65 | return (bucketized_features[split.feature_id()](index_in_batch) <= |
| 66 | split.threshold()) |
| 67 | ? split.left_id() |
| 68 | : split.right_id(); |
| 69 | } |
| 70 | case boosted_trees::Node::kCategoricalSplit: { |
| 71 | const auto& split = node.categorical_split(); |
| 72 | return (bucketized_features[split.feature_id()](index_in_batch) == |
| 73 | split.value()) |
| 74 | ? split.left_id() |
| 75 | : split.right_id(); |
| 76 | } |
| 77 | default: |
| 78 | DCHECK(false) << "Node type " << node.node_case() << " not supported."; |
| 79 | } |
| 80 | return -1; |
| 81 | } |
| 82 | |
| 83 | std::vector<float> BoostedTreesEnsembleResource::node_value( |
| 84 | const int32 tree_id, const int32 node_id) const { |
no test coverage detected