diff --git a/README.md b/README.md
index b3f92c6..9cb618b 100644
--- a/README.md
+++ b/README.md
@@ -4,7 +4,7 @@
`sgtlearn` is a Python package for learning [Shape Generalized Trees (SGTs)](https://neurips.cc/virtual/2025/loc/san-diego/poster/115950).
-- 🌳 **Shape Generalized Trees (SGTs):** A class of decision trees where each node applies a learnable, axis-aligned shape function to a feature for non-linear and interpretable splits.
+- 🌳 **Shape Generalized Trees (SGTs):** A class of decision trees where each node applies a learnable, axis-aligned shape function to one or two logical features for non-linear and interpretable splits.
- 👁 **Interpretability:** Each node's shape function can be visualized directly.
- ⚡ **ShapeCART Algorithm:** An efficient induction method for learning SGTs from data.
- 🔀 **Extensions:**
@@ -12,10 +12,9 @@
- **SGTK:** Multi-way branching generalization.
- **Shape²CART & ShapeCARTK:** Algorithms for learning S²GTs and SGTKs.
+
> [!NOTE]
-> This codebase is an efficient, but working implementation of the algorithms in the paper "Empowering Decision Trees via Shape Function Branching". Please refer to the [ROADMAP](ROADMAP.md) for a detailed list of features that are currently implemented and those that are planned for future releases. For the canonical code base for the paper, please refer to https://github.com/optimal-uoft/Empowering-DTs-via-Shape-Functions. Features in the paper that are not yet implemented in this codebase include:
-> * Bivariate shape functions (Shape$^2$CART) + Higher branching factors for bivariate splits (Shape$^2$SGT$_K$)
-> * Visualization for bivariate splits (ex. contour plots)
+> This codebase is an efficient implementation of the algorithms in "Empowering Decision Trees via Shape Function Branching." See the [ROADMAP](ROADMAP.md) for implementation status and the [canonical research code](https://github.com/optimal-uoft/Empowering-DTs-via-Shape-Functions) for the paper's original implementation.
## Installation
diff --git a/ROADMAP.md b/ROADMAP.md
index 6bfbf81..88eddd6 100644
--- a/ROADMAP.md
+++ b/ROADMAP.md
@@ -17,6 +17,12 @@
- [x] **NaN routing at predict**: if training saw missing at that split, follow the stored direction; otherwise route to the majority child.
## v0.3.0
-- [ ] multioutput support
-- [ ] Shape$^2$CART
-- [ ] Shape$^2$CART Random Forest Ensembling
+- [x] multioutput support
+- [x] Opt-in Shape$^2$CART for SGT estimators, including continuous/categorical
+ pairs, joint missing routing, and multiway branching ([tutorial](https://sgtlearn.readthedocs.io/en/latest/tutorials/bivariate-branching.html))
+- [x] Shape$^2$CART Random Forest Ensembling
+- [x] Pair-aware TAO refinement ([#48](https://github.com/optimal-uoft/sgtlearn/issues/48))
+- [x] Shape$^2$CART routing heatmap visualization ([#28](https://github.com/optimal-uoft/sgtlearn/issues/28))
+
+See the implementation specification in [#42](https://github.com/optimal-uoft/sgtlearn/issues/42)
+and the umbrella issue [#27](https://github.com/optimal-uoft/sgtlearn/issues/27).
diff --git a/cpp/CMakeLists.txt b/cpp/CMakeLists.txt
index bb50f40..7727d72 100644
--- a/cpp/CMakeLists.txt
+++ b/cpp/CMakeLists.txt
@@ -77,6 +77,10 @@ add_library(sgtlearn_core STATIC
src/Splitters/categorical/CategoricalRegressionSplitter.h
src/Splitters/categorical/CategoricalRegressionSplitter.cpp
src/Discretizers/univariate/UnivariateClassificationDiscretizer.cpp
+ src/Discretizers/pair/PairClassificationDiscretizer.h
+ src/Discretizers/pair/PairClassificationDiscretizer.cpp
+ src/Discretizers/pair/PairRegressionDiscretizer.h
+ src/Discretizers/pair/PairRegressionDiscretizer.cpp
src/Splitters/univariate/SquaredErrorSplitter.h
src/Splitters/univariate/SquaredErrorSplitter.cpp
src/Discretizers/univariate/UnivariateRegressionDiscretizer.cpp
@@ -250,4 +254,4 @@ if (SGTLEARN_BUILD_TESTS)
endif ()
-# endregion
\ No newline at end of file
+# endregion
diff --git a/cpp/bindings/ShapeGeneralizedTrees.cpp b/cpp/bindings/ShapeGeneralizedTrees.cpp
index 15ce024..52daccd 100644
--- a/cpp/bindings/ShapeGeneralizedTrees.cpp
+++ b/cpp/bindings/ShapeGeneralizedTrees.cpp
@@ -52,7 +52,8 @@ PYBIND11_MODULE(ShapeGeneralizedTrees, m) {
size_t coordinate_descent_max_iters,
size_t coordinate_descent_patience,
bool coordinate_descent_smart_init, uint64_t random_state,
- py::object max_features) {
+ py::object max_features, size_t pairwise_candidates,
+ double pairwise_penalty) {
return ClassificationShapeGeneralizedTreePy(
std::move(criterion), std::move(num_classes), num_partitions,
outer_min_leaf_size, outer_min_gain_split, outer_max_depth,
@@ -60,7 +61,8 @@ PYBIND11_MODULE(ShapeGeneralizedTrees, m) {
inner_max_depth, inner_max_leaf_nodes,
coordinate_descent_max_iters, coordinate_descent_patience,
coordinate_descent_smart_init, random_state,
- std::move(max_features));
+ std::move(max_features), pairwise_candidates,
+ pairwise_penalty);
}),
py::arg("criterion") = "gini", py::arg("num_classes"),
py::arg("num_partitions") = 2,
@@ -76,7 +78,9 @@ PYBIND11_MODULE(ShapeGeneralizedTrees, m) {
py::arg("coordinate_descent_patience") = 5,
py::arg("coordinate_descent_smart_init") = true,
py::arg("random_state") = 42,
- py::arg("max_features") = py::none())
+ py::arg("max_features") = py::none(),
+ py::arg("pairwise_candidates") = 0,
+ py::arg("pairwise_penalty") = 0.0)
.def("fit", &ClassificationShapeGeneralizedTreePy::fit, py::arg("X"),
py::arg("y"), py::arg("sample_weight") = py::none(),
py::arg("features"),
@@ -104,6 +108,8 @@ PYBIND11_MODULE(ShapeGeneralizedTrees, m) {
&ClassificationShapeGeneralizedTreePy::classesPerOutput)
.def_property_readonly(
"is_fitted", &ClassificationShapeGeneralizedTreePy::isFitted)
+ .def_property_readonly(
+ "has_pair_nodes", &ClassificationShapeGeneralizedTreePy::hasPairNodes)
.def_property_readonly(
"feature_importance",
&ClassificationShapeGeneralizedTreePy::featureImportance,
@@ -122,14 +128,16 @@ PYBIND11_MODULE(ShapeGeneralizedTrees, m) {
size_t coordinate_descent_max_iters,
size_t coordinate_descent_patience,
bool coordinate_descent_smart_init, uint64_t random_state,
- py::object max_features) {
+ py::object max_features, size_t pairwise_candidates,
+ double pairwise_penalty) {
return RegressionShapeGeneralizedTreePy(
std::move(criterion), num_partitions, outer_min_leaf_size,
outer_min_gain_split, outer_max_depth, outer_max_leaf_nodes,
inner_min_leaf_size, inner_min_gain_split, inner_max_depth,
inner_max_leaf_nodes, coordinate_descent_max_iters,
coordinate_descent_patience, coordinate_descent_smart_init,
- random_state, std::move(max_features));
+ random_state, std::move(max_features), pairwise_candidates,
+ pairwise_penalty);
}),
py::arg("criterion") = "squared_error",
py::arg("num_partitions") = 2, py::arg("outer_min_leaf_size") = 1,
@@ -142,6 +150,8 @@ PYBIND11_MODULE(ShapeGeneralizedTrees, m) {
py::arg("coordinate_descent_patience") = 5,
py::arg("coordinate_descent_smart_init") = true,
py::arg("random_state") = 42, py::arg("max_features") = py::none(),
+ py::arg("pairwise_candidates") = 0,
+ py::arg("pairwise_penalty") = 0.0,
R"(Regression tree: inner bins are round-robin seeded. ``squared_error`` runs
coordinate descent and keeps the map only if branch MSE improves clearly vs the seed;
otherwise the snapshot is restored and the branch objective is rebuilt.
@@ -164,6 +174,8 @@ is accepted for API parity with ClassificationShapeGeneralizedTree but ignored.)
&RegressionShapeGeneralizedTreePy::nOutputs)
.def_property_readonly("is_fitted",
&RegressionShapeGeneralizedTreePy::isFitted)
+ .def_property_readonly(
+ "has_pair_nodes", &RegressionShapeGeneralizedTreePy::hasPairNodes)
.def_property_readonly(
"feature_importance",
&RegressionShapeGeneralizedTreePy::featureImportance,
diff --git a/cpp/bindings/TreeAlternatingOptimization.cpp b/cpp/bindings/TreeAlternatingOptimization.cpp
index a3d15e9..40c7044 100644
--- a/cpp/bindings/TreeAlternatingOptimization.cpp
+++ b/cpp/bindings/TreeAlternatingOptimization.cpp
@@ -23,6 +23,7 @@
#include "algorithms/TAO/TreeAlternatingOptimization.h"
#include
+#include
#include
#include
#include
@@ -130,13 +131,17 @@ TaoRunContext TaoRunContext::make(py::object tree, const py::array &X,
void TreeAlternatingOptimization(py::object tree, const py::array &X,
const py::array &y,
py::object sample_weight = py::none(),
- size_t n_runs = 10, double lambda_ = 0.0) {
+ size_t n_runs = 10, double lambda_ = 0.0,
+ double tao_pair_scale = 1.1) {
+ if (!std::isfinite(tao_pair_scale) || tao_pair_scale < 0.0)
+ throw std::invalid_argument(
+ "tao_pair_scale must be finite and non-negative");
TaoRunContext ctx = TaoRunContext::make(tree, X, y, sample_weight);
if (!ctx.isFitted())
throw std::logic_error("TreeAlternatingOptimization: model is not fitted");
py::gil_scoped_release release;
- tao::optimize(ctx.adapter(), n_runs, lambda_);
+ tao::optimize(ctx.adapter(), n_runs, lambda_, tao_pair_scale);
}
} // namespace
@@ -159,11 +164,13 @@ PYBIND11_MODULE(TreeAlternatingOptimization, m) {
m.def("TreeAlternatingOptimization", &TreeAlternatingOptimization,
py::arg("tree"), py::arg("X"), py::arg("y"),
py::arg("sample_weight") = py::none(), py::arg("n_runs") = 10,
- py::arg("lambda_") = 0.0,
+ py::arg("lambda_") = 0.0, py::arg("tao_pair_scale") = 1.1,
"Refine a fitted ClassificationShapeGeneralizedTree or "
"RegressionShapeGeneralizedTree in place. X is "
"(n_samples, n_features) float32; y is 1-D class labels (uint) or "
"float targets matching the tree type. Runs up to n_runs bottom-up "
- "sweeps; lambda_ penalizes non-constant routing splits by "
- "lambda_ * totalSampleWeight in weighted reward units.");
+ "sweeps. In weighted reward units, lambda_ penalizes single-feature "
+ "routers by lambda_ * nodeSampleCount and pair routers by "
+ "tao_pair_scale * lambda_ * nodeSampleCount; dummy routers are "
+ "unpenalized.");
}
diff --git a/cpp/bindings/_sgt_estimators.h b/cpp/bindings/_sgt_estimators.h
index 1822747..d361cf6 100644
--- a/cpp/bindings/_sgt_estimators.h
+++ b/cpp/bindings/_sgt_estimators.h
@@ -17,6 +17,8 @@
#include "Discretizers/categorical/CategoricalClassificationDiscretizer.h"
#include "Discretizers/categorical/CategoricalRegressionDiscretizer.h"
+#include "Discretizers/pair/PairClassificationDiscretizer.h"
+#include "Discretizers/pair/PairRegressionDiscretizer.h"
#include "Domain/LearningCriterion.h"
#include "Domain/FeatureInfo.h"
#include "Discretizers/univariate/UnivariateDiscretizer.h"
@@ -45,6 +47,39 @@ inline py::list routingFeaturesPy(const ShapeFunctionNode &n) {
return feats;
}
+inline py::list logicalFeaturesPy(const ShapeFunctionNode &n) {
+ py::list feats;
+ for (size_t f : n.logicalFeatureIndices)
+ feats.append(f);
+ return feats;
+}
+
+inline py::list pairAxesPy(const ShapeFunctionNode &n,
+ const std::array &axes) {
+ py::list out;
+ for (size_t axis = 0; axis < axes.size(); ++axis) {
+ py::dict item;
+ item["logical_feature"] = n.logicalFeatureIndices.at(axis);
+ item["kind"] = axes[axis].type == FeatureType::Categorical
+ ? "categorical"
+ : "continuous";
+ py::list columns;
+ py::list categories;
+ for (size_t raw : axes[axis].indices) {
+ columns.append(raw);
+ if (axes[axis].type == FeatureType::Categorical)
+ categories.append(raw);
+ }
+ item["columns"] = columns;
+ item["categories"] = categories;
+ item["catchall"] = axes[axis].type == FeatureType::Categorical
+ ? py::cast("missing")
+ : py::none();
+ out.append(item);
+ }
+ return out;
+}
+
inline py::object primaryRoutingFeaturePy(const ShapeFunctionNode &n) {
if (n.routingFeatures.empty())
return py::none();
@@ -258,7 +293,8 @@ class ClassificationShapeGeneralizedTreePy {
double innerMinGainSplit, size_t innerMaxDepth, size_t innerMaxLeafNodes,
size_t coordinateDescentMaxIters, size_t coordinateDescentPatience,
bool coordinateDescentSmartInit, uint64_t random_state,
- py::object max_features = py::none()) {
+ py::object max_features = py::none(), size_t pairwiseCandidates = 0,
+ double pairwisePenalty = 0.0) {
criterionStr_ = criterion;
const LearningCriterion crit = parseClassificationCriterion(criterion);
const TreeBuildingParams outer{outerMinLeafSize, outerMinGainSplit,
@@ -271,7 +307,8 @@ class ClassificationShapeGeneralizedTreePy {
cd.smartInit = coordinateDescentSmartInit;
impl_ = std::make_unique(
crit, parseNumClassesPy(numClasses), numPartitions, outer, inner, cd,
- random_state, parseMaxFeaturesPy(max_features));
+ random_state, parseMaxFeaturesPy(max_features), pairwiseCandidates,
+ pairwisePenalty);
}
void fit(const py::array &X, const py::array &y,
@@ -334,6 +371,7 @@ class ClassificationShapeGeneralizedTreePy {
size_t numNodes() const { return impl_->numNodes(); }
bool isFitted() const { return impl_->isFitted(); }
size_t nOutputs() const { return impl_->nOutputs(); }
+ bool hasPairNodes() const { return impl_->hasPairNodes(); }
py::list classesPerOutput() const {
py::list out;
@@ -401,6 +439,46 @@ class ClassificationShapeGeneralizedTreePy {
} else {
d["feature"] = primaryRoutingFeaturePy(n);
d["features"] = routingFeaturesPy(n);
+ const auto *pairDisc = dynamic_cast(
+ n.innerDiscretizer.get());
+ if (pairDisc) {
+ d["routing_kind"] = "pair";
+ d["pair_features"] = logicalFeaturesPy(n);
+ d["pair_axes"] = pairAxesPy(n, pairDisc->axes());
+ py::list innerTree;
+ py::list leafBins;
+ const auto &routingTree = pairDisc->routingTree();
+ for (size_t nodeIndex = 0; nodeIndex < routingTree.size(); ++nodeIndex) {
+ const PairRoutingTreeNode &inner = routingTree[nodeIndex];
+ py::dict innerNode;
+ innerNode["id"] = nodeIndex;
+ innerNode["is_leaf"] = inner.isLeaf;
+ if (inner.isLeaf) {
+ innerNode["bin"] = inner.bin;
+ leafBins.append(inner.bin);
+ } else {
+ innerNode["feature"] = inner.rawFeature;
+ innerNode["axis"] = inner.featurePosition;
+ innerNode["kind"] = inner.featureType == FeatureType::Categorical
+ ? "categorical"
+ : "continuous";
+ innerNode["threshold"] =
+ inner.featureType == FeatureType::Continuous
+ ? py::cast(inner.threshold)
+ : py::none();
+ innerNode["category"] =
+ inner.featureType == FeatureType::Categorical
+ ? py::cast(inner.rawFeature)
+ : py::none();
+ innerNode["left"] = inner.left;
+ innerNode["right"] = inner.right;
+ innerNode["missing"] = inner.missing;
+ }
+ innerTree.append(innerNode);
+ }
+ d["pair_inner_tree"] = innerTree;
+ d["pair_leaf_bins"] = leafBins;
+ }
const bool isCategorical =
isCategoricalInnerDiscretizer(*n.innerDiscretizer);
d["is_categorical"] = isCategorical;
@@ -427,8 +505,9 @@ class ClassificationShapeGeneralizedTreePy {
py::list bsc;
for (size_t v : n.binSampleCounts) bsc.append(v);
d["bin_sample_counts"] = bsc;
- d["nan_prediction_partition"] =
- n.binToPartition.empty() ? 0 : n.binToPartition.back();
+ if (!pairDisc)
+ d["nan_prediction_partition"] =
+ n.binToPartition.empty() ? 0 : n.binToPartition.back();
py::list ch;
if (i < childIdx.size()) {
for (size_t c : childIdx[i]) ch.append(c);
@@ -454,7 +533,8 @@ class RegressionShapeGeneralizedTreePy {
size_t innerMinLeafSize, double innerMinGainSplit, size_t innerMaxDepth,
size_t innerMaxLeafNodes, size_t coordinateDescentMaxIters,
size_t coordinateDescentPatience, bool coordinateDescentSmartInit,
- uint64_t random_state, py::object max_features = py::none()) {
+ uint64_t random_state, py::object max_features = py::none(),
+ size_t pairwiseCandidates = 0, double pairwisePenalty = 0.0) {
criterionStr_ = criterion;
const LearningCriterion crit = parseRegressionCriterion(criterion);
const TreeBuildingParams outer{outerMinLeafSize, outerMinGainSplit,
@@ -467,7 +547,7 @@ class RegressionShapeGeneralizedTreePy {
cd.smartInit = coordinateDescentSmartInit;
impl_ = std::make_unique(
crit, numPartitions, outer, inner, cd, random_state,
- parseMaxFeaturesPy(max_features));
+ parseMaxFeaturesPy(max_features), pairwiseCandidates, pairwisePenalty);
}
void fit(const py::array &X, const py::array &y,
@@ -507,6 +587,7 @@ class RegressionShapeGeneralizedTreePy {
size_t numNodes() const { return impl_->numNodes(); }
bool isFitted() const { return impl_->isFitted(); }
size_t nOutputs() const { return impl_->nOutputs(); }
+ bool hasPairNodes() const { return impl_->hasPairNodes(); }
py::array_t featureImportance() const {
return colToNumpy(impl_->featureImportance());
@@ -570,6 +651,51 @@ class RegressionShapeGeneralizedTreePy {
d["n_samples"] = total;
d["feature"] = primaryRoutingFeaturePy(n);
d["features"] = routingFeaturesPy(n);
+ const auto *pairDisc = dynamic_cast(
+ n.innerDiscretizer.get());
+ const auto *taoPairDisc =
+ dynamic_cast(
+ n.innerDiscretizer.get());
+ if (pairDisc || taoPairDisc) {
+ d["routing_kind"] = "pair";
+ d["pair_features"] = logicalFeaturesPy(n);
+ d["pair_axes"] = pairAxesPy(
+ n, pairDisc ? pairDisc->axes() : taoPairDisc->axes());
+ py::list innerTree;
+ py::list leafBins;
+ const auto &routingTree =
+ pairDisc ? pairDisc->routingTree() : taoPairDisc->routingTree();
+ for (size_t nodeIndex = 0; nodeIndex < routingTree.size(); ++nodeIndex) {
+ const PairRoutingTreeNode &inner = routingTree[nodeIndex];
+ py::dict innerNode;
+ innerNode["id"] = nodeIndex;
+ innerNode["is_leaf"] = inner.isLeaf;
+ if (inner.isLeaf) {
+ innerNode["bin"] = inner.bin;
+ leafBins.append(inner.bin);
+ } else {
+ innerNode["feature"] = inner.rawFeature;
+ innerNode["axis"] = inner.featurePosition;
+ innerNode["kind"] = inner.featureType == FeatureType::Categorical
+ ? "categorical"
+ : "continuous";
+ innerNode["threshold"] =
+ inner.featureType == FeatureType::Continuous
+ ? py::cast(inner.threshold)
+ : py::none();
+ innerNode["category"] =
+ inner.featureType == FeatureType::Categorical
+ ? py::cast(inner.rawFeature)
+ : py::none();
+ innerNode["left"] = inner.left;
+ innerNode["right"] = inner.right;
+ innerNode["missing"] = inner.missing;
+ }
+ innerTree.append(innerNode);
+ }
+ d["pair_inner_tree"] = innerTree;
+ d["pair_leaf_bins"] = leafBins;
+ }
const bool isCategorical =
isCategoricalInnerDiscretizer(*n.innerDiscretizer);
d["is_categorical"] = isCategorical;
@@ -596,8 +722,9 @@ class RegressionShapeGeneralizedTreePy {
py::list bsc;
for (size_t v : n.binSampleCounts) bsc.append(v);
d["bin_sample_counts"] = bsc;
- d["nan_prediction_partition"] =
- n.binToPartition.empty() ? 0 : n.binToPartition.back();
+ if (!pairDisc && !taoPairDisc)
+ d["nan_prediction_partition"] =
+ n.binToPartition.empty() ? 0 : n.binToPartition.back();
py::list ch;
if (i < childIdx.size()) {
for (size_t c : childIdx[i]) ch.append(c);
diff --git a/cpp/src/Discretizers/pair/PairClassificationDiscretizer.cpp b/cpp/src/Discretizers/pair/PairClassificationDiscretizer.cpp
new file mode 100644
index 0000000..50268d8
--- /dev/null
+++ b/cpp/src/Discretizers/pair/PairClassificationDiscretizer.cpp
@@ -0,0 +1,351 @@
+#include "Discretizers/pair/PairClassificationDiscretizer.h"
+
+#include "Criterion.h"
+#include "algorithms/missing_values.h"
+
+#include
+#include
+#include
+#include
+#include
+#include
+
+PairClassificationDiscretizer::PairClassificationDiscretizer(
+ LearningCriterion criterion, FeatureInfo first, FeatureInfo second)
+ : criterion_(criterion), axes_{std::move(first), std::move(second)} {
+ axisOffsets_[1] = axes_[0].indices.n_elem;
+ axisOffsets_[2] = axisOffsets_[1] + axes_[1].indices.n_elem;
+ if (axisOffsets_[1] == 0 || axisOffsets_[2] == axisOffsets_[1])
+ throw std::invalid_argument("pair CART logical features cannot be empty");
+ for (const FeatureInfo &axis : axes_)
+ if (axis.type == FeatureType::Continuous && axis.indices.n_elem != 1)
+ throw std::invalid_argument(
+ "pair CART continuous logical features require one column");
+ routingFeatures_ = arma::join_cols(axes_[0].indices, axes_[1].indices);
+}
+
+bool PairClassificationDiscretizer::axisMissing(
+ size_t axis, const arma::fmat &X, size_t sample) const {
+ if (axes_[axis].type == FeatureType::Continuous)
+ return !missing_values::is_finite(X(axes_[axis].indices(0), sample));
+ for (size_t raw : axes_[axis].indices)
+ if (X(raw, sample) >= 0.5F)
+ return false;
+ return true;
+}
+
+bool PairClassificationDiscretizer::axisMissing(
+ size_t axis, const std::vector &values) const {
+ if (axes_[axis].type == FeatureType::Continuous)
+ return !missing_values::is_finite(values[axisOffsets_[axis]]);
+ for (size_t pos = axisOffsets_[axis]; pos < axisOffsets_[axis + 1]; ++pos)
+ if (values[pos] >= 0.5F)
+ return false;
+ return true;
+}
+
+size_t PairClassificationDiscretizer::routingPosition(
+ size_t rawFeature) const {
+ const auto it =
+ std::find(routingFeatures_.begin(), routingFeatures_.end(), rawFeature);
+ if (it == routingFeatures_.end())
+ throw std::runtime_error("pair CART routing feature not found");
+ return static_cast(it - routingFeatures_.begin());
+}
+
+double PairClassificationDiscretizer::impurity(
+ const std::vector> &stats) const {
+ if (criterion_ == LearningCriterion::Gini)
+ return Criterion::gini(stats);
+ if (criterion_ == LearningCriterion::Entropy)
+ return Criterion::entropy(stats);
+ throw std::invalid_argument("pair CART requires a classification criterion");
+}
+
+PairClassificationDiscretizer::Split
+PairClassificationDiscretizer::bestSplit(
+ const BuildNode &node, const arma::fmat &X, const arma::Mat &y,
+ const std::vector &classes, const arma::Row &weights,
+ size_t minLeafSize, double totalWeight) const {
+ Split best;
+ double bestChildScore = std::numeric_limits::infinity();
+
+ const auto makeStats = [&y, &classes, &weights](
+ const std::vector &samples) {
+ std::vector> stats(classes.size());
+ for (size_t o = 0; o < classes.size(); ++o)
+ stats[o].assign(classes[o], 0.0);
+ for (size_t sample : samples)
+ for (size_t o = 0; o < classes.size(); ++o)
+ stats[o][y(o, sample)] += weights(sample);
+ return stats;
+ };
+ const auto sumWeight = [&weights](const std::vector &samples) {
+ double total = 0.0;
+ for (size_t sample : samples)
+ total += weights(sample);
+ return total;
+ };
+
+ for (size_t axis = 0; axis < 2; ++axis) {
+ std::vector valid;
+ std::vector missing;
+ for (size_t sample : node.samples)
+ (axisMissing(axis, X, sample) ? missing : valid).push_back(sample);
+
+ const auto missingStats = makeStats(missing);
+ const double missingWeight = sumWeight(missing);
+ const auto consider = [&](size_t rawFeature, double threshold,
+ std::vector leftSamples,
+ std::vector rightSamples,
+ const std::vector> &leftStats,
+ const std::vector> &rightStats,
+ double leftWeight, double rightWeight) {
+ if (leftSamples.size() < minLeafSize ||
+ rightSamples.size() < minLeafSize)
+ return;
+ const double childScore =
+ node.weight > 0.0
+ ? (leftWeight * impurity(leftStats) +
+ rightWeight * impurity(rightStats) +
+ missingWeight * impurity(missingStats)) /
+ node.weight
+ : 0.0;
+ if (childScore >=
+ bestChildScore - std::numeric_limits::epsilon())
+ return;
+ best.found = true;
+ bestChildScore = childScore;
+ best.featurePosition = axis;
+ best.rawFeature = rawFeature;
+ best.featureType = axes_[axis].type;
+ best.threshold = threshold;
+ best.left = std::move(leftSamples);
+ best.right = std::move(rightSamples);
+ best.missing = missing;
+ best.gain = totalWeight > 0.0
+ ? (node.weight / totalWeight) *
+ (node.impurity - childScore)
+ : 0.0;
+ };
+
+ if (axes_[axis].type == FeatureType::Categorical) {
+ for (size_t rawFeature : axes_[axis].indices) {
+ std::vector left;
+ std::vector right;
+ for (size_t sample : valid)
+ (X(rawFeature, sample) >= 0.5F ? right : left).push_back(sample);
+ consider(rawFeature, 0.5, left, right, makeStats(left),
+ makeStats(right), sumWeight(left), sumWeight(right));
+ }
+ continue;
+ }
+
+ const size_t rawFeature = axes_[axis].indices(0);
+ std::vector order = std::move(valid);
+ std::stable_sort(order.begin(), order.end(),
+ [&X, rawFeature](size_t a, size_t b) {
+ return X(rawFeature, a) < X(rawFeature, b);
+ });
+
+ std::vector> left(classes.size());
+ std::vector> right = makeStats(order);
+ for (size_t o = 0; o < classes.size(); ++o)
+ left[o].assign(classes[o], 0.0);
+ double leftWeight = 0.0;
+ const double validWeight = node.weight - missingWeight;
+
+ for (size_t i = 1; i < order.size(); ++i) {
+ const size_t moved = order[i - 1];
+ const double w = weights(moved);
+ leftWeight += w;
+ for (size_t o = 0; o < classes.size(); ++o) {
+ left[o][y(o, moved)] += w;
+ right[o][y(o, moved)] -= w;
+ }
+ if (i < minLeafSize)
+ continue;
+ if (order.size() - i < minLeafSize)
+ break;
+
+ const float previous = X(rawFeature, order[i - 1]);
+ const float current = X(rawFeature, order[i]);
+ if (current <= previous + 1e-7F)
+ continue;
+ double threshold = static_cast(previous) / 2.0 +
+ static_cast(current) / 2.0;
+ if (!std::isfinite(threshold) || threshold == current)
+ threshold = previous;
+ consider(
+ rawFeature, threshold,
+ std::vector(order.begin(), order.begin() + i),
+ std::vector(order.begin() + i, order.end()), left, right,
+ leftWeight, validWeight - leftWeight);
+ }
+ }
+ return best;
+}
+
+void PairClassificationDiscretizer::Train(
+ const arma::fmat &X, arma::uvec &features, const arma::Mat &y,
+ const std::vector &classes, size_t minLeafSize,
+ double minGainSplit, size_t maxDepth, size_t maxLeafNodes,
+ const arma::Row &sampleWeights) {
+ this->resetTrainedOutputs();
+ tree_.clear();
+ if (features.n_elem != routingFeatures_.n_elem || y.n_cols != X.n_cols ||
+ classes.size() != y.n_rows)
+ throw std::invalid_argument("invalid pair CART training shapes");
+
+ arma::Row weights = sampleWeights;
+ if (weights.n_elem == 0)
+ weights.ones(X.n_cols);
+ if (weights.n_elem != X.n_cols)
+ throw std::invalid_argument("pair CART sample weight length mismatch");
+ for (size_t rawFeature : routingFeatures_)
+ if (rawFeature >= X.n_rows)
+ throw std::invalid_argument("pair CART feature out of range");
+
+ const auto makeNode = [&y, &classes, &weights, this](
+ std::vector samples, size_t depth) {
+ BuildNode node;
+ node.samples = std::move(samples);
+ node.depth = depth;
+ node.stats.resize(classes.size());
+ for (size_t o = 0; o < classes.size(); ++o)
+ node.stats[o].assign(classes[o], 0.0);
+ for (size_t sample : node.samples) {
+ const double w = weights(sample);
+ node.weight += w;
+ for (size_t o = 0; o < classes.size(); ++o)
+ node.stats[o][y(o, sample)] += w;
+ }
+ node.impurity = impurity(node.stats);
+ return node;
+ };
+
+ std::vector all(X.n_cols);
+ std::iota(all.begin(), all.end(), 0);
+ std::vector nodes;
+ nodes.push_back(makeNode(std::move(all), 0));
+ const double totalWeight = nodes.front().weight;
+ size_t finiteLeafCount = 1;
+
+ const auto splitNode = [&](size_t index, Split split) {
+ const size_t depth = nodes[index].depth;
+ const bool missingGoesLeft = split.left.size() >= split.right.size();
+ nodes[index].routing.isLeaf = false;
+ nodes[index].routing.featurePosition = split.featurePosition;
+ nodes[index].routing.rawFeature = split.rawFeature;
+ nodes[index].routing.featureType = split.featureType;
+ nodes[index].routing.threshold = split.threshold;
+ nodes[index].routing.left = nodes.size();
+ nodes.push_back(makeNode(std::move(split.left), depth + 1));
+ nodes[index].routing.right = nodes.size();
+ nodes.push_back(makeNode(std::move(split.right), depth + 1));
+ if (split.missing.empty()) {
+ nodes[index].routing.missing = missingGoesLeft
+ ? nodes[index].routing.left
+ : nodes[index].routing.right;
+ } else {
+ nodes[index].routing.missing = nodes.size();
+ nodes.push_back(makeNode(std::move(split.missing), depth + 1));
+ }
+ ++finiteLeafCount;
+ };
+
+ if (maxLeafNodes == 0) {
+ const auto grow = [&](auto &&self, size_t index) -> void {
+ if (maxDepth != 0 && nodes[index].depth >= maxDepth)
+ return;
+ Split split = bestSplit(nodes[index], X, y, classes, weights,
+ minLeafSize, totalWeight);
+ if (!split.found ||
+ split.gain + std::numeric_limits::epsilon() < minGainSplit)
+ return;
+ splitNode(index, std::move(split));
+ self(self, nodes[index].routing.left);
+ self(self, nodes[index].routing.right);
+ if (nodes[index].routing.missing != nodes[index].routing.left)
+ self(self, nodes[index].routing.missing);
+ };
+ grow(grow, 0);
+ } else {
+ while (finiteLeafCount < maxLeafNodes) {
+ size_t bestIndex = nodes.size();
+ Split best;
+ for (size_t i = 0; i < nodes.size(); ++i) {
+ if (!nodes[i].routing.isLeaf ||
+ (maxDepth != 0 && nodes[i].depth >= maxDepth))
+ continue;
+ Split candidate = bestSplit(nodes[i], X, y, classes, weights,
+ minLeafSize, totalWeight);
+ if (!candidate.found ||
+ candidate.gain + std::numeric_limits::epsilon() <
+ minGainSplit)
+ continue;
+ if (bestIndex == nodes.size() || candidate.gain > best.gain) {
+ bestIndex = i;
+ best = std::move(candidate);
+ }
+ }
+ if (bestIndex == nodes.size())
+ break;
+ splitNode(bestIndex, std::move(best));
+ }
+ }
+
+ for (BuildNode &node : nodes) {
+ if (!node.routing.isLeaf)
+ continue;
+ node.routing.bin = this->inSampleDiscretizations_.size();
+ this->inSampleDiscretizations_.push_back(std::move(node.samples));
+ this->leafStats_.push_back(std::move(node.stats));
+ this->leafNumSamples_.push_back(
+ this->inSampleDiscretizations_.back().size());
+ this->leafNodeWeights_.push_back(node.weight);
+ }
+ tree_.reserve(nodes.size());
+ for (const BuildNode &node : nodes)
+ tree_.push_back(node.routing);
+ this->numLeaves_ = this->leafStats_.size();
+ this->markTrained();
+}
+
+size_t PairClassificationDiscretizer::routeValues(
+ const std::vector &values) const {
+ if (values.size() != routingFeatures_.n_elem)
+ throw std::invalid_argument(
+ "pair router values do not match routing features");
+ size_t index = 0;
+ while (!tree_[index].isLeaf) {
+ const PairRoutingTreeNode &node = tree_[index];
+ if (axisMissing(node.featurePosition, values)) {
+ index = node.missing;
+ continue;
+ }
+ const float value = values[routingPosition(node.rawFeature)];
+ index = node.featureType == FeatureType::Categorical
+ ? (value >= 0.5F ? node.right : node.left)
+ : (value <= node.threshold ? node.left : node.right);
+ }
+ return tree_[index].bin;
+}
+
+size_t PairClassificationDiscretizer::routeToBin(
+ const std::vector &featureValues) const {
+ this->ensureTrained();
+ return routeValues(featureValues);
+}
+
+void PairClassificationDiscretizer::transform(
+ const arma::fmat &X, arma::Row &binLoc) const {
+ this->ensureTrained();
+ binLoc.set_size(X.n_cols);
+ std::vector values(routingFeatures_.n_elem);
+ for (arma::uword i = 0; i < X.n_cols; ++i) {
+ for (size_t j = 0; j < routingFeatures_.n_elem; ++j)
+ values[j] = X(routingFeatures_(j), i);
+ binLoc(i) = routeValues(values);
+ }
+}
diff --git a/cpp/src/Discretizers/pair/PairClassificationDiscretizer.h b/cpp/src/Discretizers/pair/PairClassificationDiscretizer.h
new file mode 100644
index 0000000..52c361a
--- /dev/null
+++ b/cpp/src/Discretizers/pair/PairClassificationDiscretizer.h
@@ -0,0 +1,81 @@
+#pragma once
+
+#include "Discretizers/ClassificationDiscretizer.h"
+#include "Domain/FeatureInfo.h"
+#include "Domain/LearningCriterion.h"
+
+#include
+#include
+#include
+#include
+
+struct PairRoutingTreeNode {
+ bool isLeaf = true;
+ size_t rawFeature = 0;
+ size_t featurePosition = 0;
+ FeatureType featureType = FeatureType::Continuous;
+ double threshold = 0.0;
+ size_t left = 0;
+ size_t right = 0;
+ size_t missing = 0;
+ size_t bin = 0;
+};
+
+/** Ordinary axis-aligned CART over exactly two logical features. */
+class PairClassificationDiscretizer final : public ClassificationDiscretizer {
+public:
+ PairClassificationDiscretizer(LearningCriterion criterion,
+ FeatureInfo first, FeatureInfo second);
+
+ void Train(const arma::fmat &X, arma::uvec &features,
+ const arma::Mat &y,
+ const std::vector &nClassesPerOutput,
+ size_t minLeafSize, double minGainSplit, size_t maxDepth,
+ size_t maxLeafNodes,
+ const arma::Row &sampleWeights = arma::Row()) override;
+
+ void transform(const arma::fmat &X, arma::Row &binLoc) const override;
+ size_t routeToBin(const std::vector &featureValues) const override;
+
+ const std::vector &routingTree() const { return tree_; }
+ const std::array &axes() const { return axes_; }
+
+private:
+ struct BuildNode {
+ PairRoutingTreeNode routing;
+ std::vector samples;
+ std::vector> stats;
+ double weight = 0.0;
+ double impurity = 0.0;
+ size_t depth = 0;
+ };
+
+ struct Split {
+ bool found = false;
+ size_t featurePosition = 0;
+ size_t rawFeature = 0;
+ FeatureType featureType = FeatureType::Continuous;
+ double threshold = 0.0;
+ double gain = 0.0;
+ std::vector left;
+ std::vector right;
+ std::vector missing;
+ };
+
+ double impurity(const std::vector> &stats) const;
+ Split bestSplit(const BuildNode &node, const arma::fmat &X,
+ const arma::Mat &y,
+ const std::vector &classes,
+ const arma::Row &weights, size_t minLeafSize,
+ double totalWeight) const;
+ size_t routeValues(const std::vector &values) const;
+ bool axisMissing(size_t axis, const arma::fmat &X, size_t sample) const;
+ bool axisMissing(size_t axis, const std::vector &values) const;
+ size_t routingPosition(size_t rawFeature) const;
+
+ LearningCriterion criterion_;
+ std::array axes_;
+ arma::uvec routingFeatures_;
+ std::array axisOffsets_{};
+ std::vector tree_;
+};
diff --git a/cpp/src/Discretizers/pair/PairRegressionDiscretizer.cpp b/cpp/src/Discretizers/pair/PairRegressionDiscretizer.cpp
new file mode 100644
index 0000000..0e05ba6
--- /dev/null
+++ b/cpp/src/Discretizers/pair/PairRegressionDiscretizer.cpp
@@ -0,0 +1,344 @@
+#include "Discretizers/pair/PairRegressionDiscretizer.h"
+
+#include "Criterion.h"
+#include "algorithms/missing_values.h"
+
+#include
+#include
+#include
+#include
+#include
+#include
+
+PairRegressionDiscretizer::PairRegressionDiscretizer(
+ LearningCriterion criterion, FeatureInfo first, FeatureInfo second)
+ : criterion_(criterion), axes_{std::move(first), std::move(second)} {
+ axisOffsets_[1] = axes_[0].indices.n_elem;
+ axisOffsets_[2] = axisOffsets_[1] + axes_[1].indices.n_elem;
+ if (axisOffsets_[1] == 0 || axisOffsets_[2] == axisOffsets_[1])
+ throw std::invalid_argument("pair CART logical features cannot be empty");
+ for (const FeatureInfo &axis : axes_)
+ if (axis.type == FeatureType::Continuous && axis.indices.n_elem != 1)
+ throw std::invalid_argument(
+ "pair CART continuous logical features require one column");
+ routingFeatures_ = arma::join_cols(axes_[0].indices, axes_[1].indices);
+ if (criterion_ != LearningCriterion::SquaredError &&
+ criterion_ != LearningCriterion::AbsoluteError)
+ throw std::invalid_argument("pair CART requires a regression criterion");
+}
+
+bool PairRegressionDiscretizer::axisMissing(
+ size_t axis, const arma::fmat &X, size_t sample) const {
+ if (axes_[axis].type == FeatureType::Continuous)
+ return !missing_values::is_finite(X(axes_[axis].indices(0), sample));
+ for (size_t raw : axes_[axis].indices)
+ if (X(raw, sample) >= 0.5F)
+ return false;
+ return true;
+}
+
+bool PairRegressionDiscretizer::axisMissing(
+ size_t axis, const std::vector &values) const {
+ if (axes_[axis].type == FeatureType::Continuous)
+ return !missing_values::is_finite(values[axisOffsets_[axis]]);
+ for (size_t pos = axisOffsets_[axis]; pos < axisOffsets_[axis + 1]; ++pos)
+ if (values[pos] >= 0.5F)
+ return false;
+ return true;
+}
+
+size_t PairRegressionDiscretizer::routingPosition(size_t rawFeature) const {
+ const auto it =
+ std::find(routingFeatures_.begin(), routingFeatures_.end(), rawFeature);
+ if (it == routingFeatures_.end())
+ throw std::runtime_error("pair CART routing feature not found");
+ return static_cast(it - routingFeatures_.begin());
+}
+
+double PairRegressionDiscretizer::impurity(
+ const std::vector &samples, const arma::Mat &y,
+ const arma::Row &weights) const {
+ if (samples.empty())
+ return 0.0;
+ if (criterion_ == LearningCriterion::SquaredError) {
+ std::vector> stats(
+ y.n_rows, std::vector(2, 0.0));
+ double totalWeight = 0.0;
+ for (size_t sample : samples) {
+ const double w = weights(sample);
+ totalWeight += w;
+ for (arma::uword o = 0; o < y.n_rows; ++o) {
+ const double value = y(o, sample);
+ stats[o][0] += w * value;
+ stats[o][1] += w * value * value;
+ }
+ }
+ return Criterion::squaredError(stats, totalWeight);
+ }
+
+ double total = 0.0;
+ for (arma::uword o = 0; o < y.n_rows; ++o) {
+ std::vector values;
+ std::vector sampleWeights;
+ values.reserve(samples.size());
+ sampleWeights.reserve(samples.size());
+ for (size_t sample : samples) {
+ values.push_back(y(o, sample));
+ sampleWeights.push_back(weights(sample));
+ }
+ total += Criterion::absoluteError(values, sampleWeights).mae;
+ }
+ return total;
+}
+
+PairRegressionDiscretizer::Split PairRegressionDiscretizer::bestSplit(
+ const BuildNode &node, const arma::fmat &X, const arma::Mat &y,
+ const arma::Row &weights, size_t minLeafSize,
+ double totalWeight) const {
+ Split best;
+ double bestChildScore = std::numeric_limits::infinity();
+ const auto sumWeight = [&weights](const std::vector &samples) {
+ double total = 0.0;
+ for (size_t sample : samples)
+ total += weights(sample);
+ return total;
+ };
+
+ for (size_t axis = 0; axis < 2; ++axis) {
+ std::vector valid;
+ std::vector missing;
+ for (size_t sample : node.samples)
+ (axisMissing(axis, X, sample) ? missing : valid).push_back(sample);
+ const double missingWeight = sumWeight(missing);
+
+ const auto consider = [&](size_t rawFeature, double threshold,
+ std::vector left,
+ std::vector right) {
+ if (left.size() < minLeafSize || right.size() < minLeafSize)
+ return;
+ const double leftWeight = sumWeight(left);
+ const double rightWeight = sumWeight(right);
+ const double childScore =
+ node.weight > 0.0
+ ? (leftWeight * impurity(left, y, weights) +
+ rightWeight * impurity(right, y, weights) +
+ missingWeight * impurity(missing, y, weights)) /
+ node.weight
+ : 0.0;
+ if (childScore >=
+ bestChildScore - std::numeric_limits::epsilon())
+ return;
+ best.found = true;
+ bestChildScore = childScore;
+ best.featurePosition = axis;
+ best.rawFeature = rawFeature;
+ best.featureType = axes_[axis].type;
+ best.threshold = threshold;
+ best.left = std::move(left);
+ best.right = std::move(right);
+ best.missing = missing;
+ best.gain = totalWeight > 0.0
+ ? (node.weight / totalWeight) *
+ (node.impurity - childScore)
+ : 0.0;
+ };
+
+ if (axes_[axis].type == FeatureType::Categorical) {
+ for (size_t rawFeature : axes_[axis].indices) {
+ std::vector left;
+ std::vector right;
+ for (size_t sample : valid)
+ (X(rawFeature, sample) >= 0.5F ? right : left).push_back(sample);
+ consider(rawFeature, 0.5, std::move(left), std::move(right));
+ }
+ continue;
+ }
+
+ const size_t rawFeature = axes_[axis].indices(0);
+ std::vector order = std::move(valid);
+ std::stable_sort(order.begin(), order.end(),
+ [&X, rawFeature](size_t a, size_t b) {
+ return X(rawFeature, a) < X(rawFeature, b);
+ });
+ for (size_t i = minLeafSize; i + minLeafSize <= order.size(); ++i) {
+ const float previous = X(rawFeature, order[i - 1]);
+ const float current = X(rawFeature, order[i]);
+ if (current <= previous + 1e-7F)
+ continue;
+ double threshold = static_cast(previous) / 2.0 +
+ static_cast(current) / 2.0;
+ if (!std::isfinite(threshold) || threshold == current)
+ threshold = previous;
+ consider(rawFeature, threshold,
+ std::vector(order.begin(), order.begin() + i),
+ std::vector(order.begin() + i, order.end()));
+ }
+ }
+ return best;
+}
+
+void PairRegressionDiscretizer::Train(
+ const arma::fmat &X, arma::uvec &features, const arma::Mat &y,
+ size_t minLeafSize, double minGainSplit, size_t maxDepth,
+ size_t maxLeafNodes, const arma::Row &sampleWeights) {
+ this->resetTrainedOutputs();
+ tree_.clear();
+ if (features.n_elem != routingFeatures_.n_elem || y.n_cols != X.n_cols)
+ throw std::invalid_argument("invalid pair CART training shapes");
+
+ arma::Row weights = sampleWeights;
+ if (weights.n_elem == 0)
+ weights.ones(X.n_cols);
+ if (weights.n_elem != X.n_cols)
+ throw std::invalid_argument("pair CART sample weight length mismatch");
+ for (size_t rawFeature : routingFeatures_)
+ if (rawFeature >= X.n_rows)
+ throw std::invalid_argument("pair CART feature out of range");
+
+ const auto makeNode = [&y, &weights, this](std::vector samples,
+ size_t depth) {
+ BuildNode node;
+ node.samples = std::move(samples);
+ node.depth = depth;
+ node.stats.assign(y.n_rows, std::vector(2, 0.0));
+ for (size_t sample : node.samples) {
+ const double w = weights(sample);
+ node.weight += w;
+ for (arma::uword o = 0; o < y.n_rows; ++o) {
+ const double value = y(o, sample);
+ node.stats[o][0] += w * value;
+ node.stats[o][1] += w * value * value;
+ }
+ }
+ node.impurity = impurity(node.samples, y, weights);
+ return node;
+ };
+
+ std::vector all(X.n_cols);
+ std::iota(all.begin(), all.end(), 0);
+ std::vector nodes;
+ nodes.push_back(makeNode(std::move(all), 0));
+ const double totalWeight = nodes.front().weight;
+ size_t finiteLeafCount = 1;
+
+ const auto splitNode = [&](size_t index, Split split) {
+ const size_t depth = nodes[index].depth;
+ const bool missingGoesLeft = split.left.size() >= split.right.size();
+ nodes[index].routing.isLeaf = false;
+ nodes[index].routing.featurePosition = split.featurePosition;
+ nodes[index].routing.rawFeature = split.rawFeature;
+ nodes[index].routing.featureType = split.featureType;
+ nodes[index].routing.threshold = split.threshold;
+ nodes[index].routing.left = nodes.size();
+ nodes.push_back(makeNode(std::move(split.left), depth + 1));
+ nodes[index].routing.right = nodes.size();
+ nodes.push_back(makeNode(std::move(split.right), depth + 1));
+ if (split.missing.empty()) {
+ nodes[index].routing.missing = missingGoesLeft
+ ? nodes[index].routing.left
+ : nodes[index].routing.right;
+ } else {
+ nodes[index].routing.missing = nodes.size();
+ nodes.push_back(makeNode(std::move(split.missing), depth + 1));
+ }
+ ++finiteLeafCount;
+ };
+
+ if (maxLeafNodes == 0) {
+ const auto grow = [&](auto &&self, size_t index) -> void {
+ if (maxDepth != 0 && nodes[index].depth >= maxDepth)
+ return;
+ Split split =
+ bestSplit(nodes[index], X, y, weights, minLeafSize, totalWeight);
+ if (!split.found ||
+ split.gain + std::numeric_limits::epsilon() < minGainSplit)
+ return;
+ splitNode(index, std::move(split));
+ self(self, nodes[index].routing.left);
+ self(self, nodes[index].routing.right);
+ if (nodes[index].routing.missing != nodes[index].routing.left)
+ self(self, nodes[index].routing.missing);
+ };
+ grow(grow, 0);
+ } else {
+ while (finiteLeafCount < maxLeafNodes) {
+ size_t bestIndex = nodes.size();
+ Split best;
+ for (size_t i = 0; i < nodes.size(); ++i) {
+ if (!nodes[i].routing.isLeaf ||
+ (maxDepth != 0 && nodes[i].depth >= maxDepth))
+ continue;
+ Split candidate =
+ bestSplit(nodes[i], X, y, weights, minLeafSize, totalWeight);
+ if (!candidate.found ||
+ candidate.gain + std::numeric_limits::epsilon() <
+ minGainSplit)
+ continue;
+ if (bestIndex == nodes.size() || candidate.gain > best.gain) {
+ bestIndex = i;
+ best = std::move(candidate);
+ }
+ }
+ if (bestIndex == nodes.size())
+ break;
+ splitNode(bestIndex, std::move(best));
+ }
+ }
+
+ for (BuildNode &node : nodes) {
+ if (!node.routing.isLeaf)
+ continue;
+ node.routing.bin = this->inSampleDiscretizations_.size();
+ this->inSampleDiscretizations_.push_back(std::move(node.samples));
+ this->leafStats_.push_back(
+ criterion_ == LearningCriterion::SquaredError
+ ? std::move(node.stats)
+ : std::vector>{});
+ this->leafNumSamples_.push_back(
+ this->inSampleDiscretizations_.back().size());
+ this->leafNodeWeights_.push_back(node.weight);
+ }
+ tree_.reserve(nodes.size());
+ for (const BuildNode &node : nodes)
+ tree_.push_back(node.routing);
+ this->numLeaves_ = this->leafStats_.size();
+ this->markTrained();
+}
+
+size_t PairRegressionDiscretizer::routeValues(
+ const std::vector &values) const {
+ if (values.size() != routingFeatures_.n_elem)
+ throw std::invalid_argument(
+ "pair router values do not match routing features");
+ size_t index = 0;
+ while (!tree_[index].isLeaf) {
+ const PairRoutingTreeNode &node = tree_[index];
+ if (axisMissing(node.featurePosition, values)) {
+ index = node.missing;
+ continue;
+ }
+ const float value = values[routingPosition(node.rawFeature)];
+ index = node.featureType == FeatureType::Categorical
+ ? (value >= 0.5F ? node.right : node.left)
+ : (value <= node.threshold ? node.left : node.right);
+ }
+ return tree_[index].bin;
+}
+
+size_t PairRegressionDiscretizer::routeToBin(
+ const std::vector &featureValues) const {
+ this->ensureTrained();
+ return routeValues(featureValues);
+}
+
+void PairRegressionDiscretizer::transform(const arma::fmat &X,
+ arma::Row &binLoc) const {
+ this->ensureTrained();
+ binLoc.set_size(X.n_cols);
+ std::vector values(routingFeatures_.n_elem);
+ for (arma::uword i = 0; i < X.n_cols; ++i) {
+ for (size_t j = 0; j < routingFeatures_.n_elem; ++j)
+ values[j] = X(routingFeatures_(j), i);
+ binLoc(i) = routeValues(values);
+ }
+}
diff --git a/cpp/src/Discretizers/pair/PairRegressionDiscretizer.h b/cpp/src/Discretizers/pair/PairRegressionDiscretizer.h
new file mode 100644
index 0000000..0f3be9d
--- /dev/null
+++ b/cpp/src/Discretizers/pair/PairRegressionDiscretizer.h
@@ -0,0 +1,68 @@
+#pragma once
+
+#include "Discretizers/RegressionDiscretizer.h"
+#include "Discretizers/pair/PairClassificationDiscretizer.h"
+#include "Domain/LearningCriterion.h"
+
+#include
+#include
+#include
+#include
+
+/** Ordinary axis-aligned CART over exactly two logical features. */
+class PairRegressionDiscretizer final : public RegressionDiscretizer {
+public:
+ PairRegressionDiscretizer(LearningCriterion criterion, FeatureInfo first,
+ FeatureInfo second);
+
+ void Train(const arma::fmat &X, arma::uvec &features,
+ const arma::Mat &y, size_t minLeafSize,
+ double minGainSplit, size_t maxDepth, size_t maxLeafNodes,
+ const arma::Row &sampleWeights = arma::Row()) override;
+
+ void transform(const arma::fmat &X, arma::Row &binLoc) const override;
+ size_t routeToBin(const std::vector &featureValues) const override;
+
+ const std::vector &routingTree() const { return tree_; }
+ const std::array &axes() const { return axes_; }
+
+private:
+ struct BuildNode {
+ PairRoutingTreeNode routing;
+ std::vector samples;
+ std::vector> stats;
+ double weight = 0.0;
+ double impurity = 0.0;
+ size_t depth = 0;
+ };
+
+ struct Split {
+ bool found = false;
+ size_t featurePosition = 0;
+ size_t rawFeature = 0;
+ FeatureType featureType = FeatureType::Continuous;
+ double threshold = 0.0;
+ double gain = 0.0;
+ std::vector left;
+ std::vector right;
+ std::vector missing;
+ };
+
+ double impurity(const std::vector &samples,
+ const arma::Mat &y,
+ const arma::Row &weights) const;
+ Split bestSplit(const BuildNode &node, const arma::fmat &X,
+ const arma::Mat &y,
+ const arma::Row &weights, size_t minLeafSize,
+ double totalWeight) const;
+ size_t routeValues(const std::vector &values) const;
+ bool axisMissing(size_t axis, const arma::fmat &X, size_t sample) const;
+ bool axisMissing(size_t axis, const std::vector &values) const;
+ size_t routingPosition(size_t rawFeature) const;
+
+ LearningCriterion criterion_;
+ std::array axes_;
+ arma::uvec routingFeatures_;
+ std::array axisOffsets_{};
+ std::vector tree_;
+};
diff --git a/cpp/src/Estimators/ClassificationShapeGeneralizedTree.cpp b/cpp/src/Estimators/ClassificationShapeGeneralizedTree.cpp
index 2c24472..1f10e89 100644
--- a/cpp/src/Estimators/ClassificationShapeGeneralizedTree.cpp
+++ b/cpp/src/Estimators/ClassificationShapeGeneralizedTree.cpp
@@ -14,6 +14,7 @@
#include "Criterion.h"
#include "Discretizers/ClassificationDiscretizer.h"
+#include "Discretizers/pair/PairClassificationDiscretizer.h"
#include "Discretizers/factories/DiscretizerFactories.h"
#include "Estimators/ShapeFunctions/ShapeFunctionSplitSearch.h"
@@ -47,13 +48,15 @@ ClassificationShapeGeneralizedTree::ClassificationShapeGeneralizedTree(
LearningCriterion criterion, std::vector numClasses,
size_t numPartitions, TreeBuildingParams outerParams,
TreeBuildingParams innerParams, CoordinateDescentParams cdParams,
- uint64_t random_state, FeatureBaggingPickFn featureBagging)
+ uint64_t random_state, FeatureBaggingPickFn featureBagging,
+ size_t pairwiseCandidates, double pairwisePenalty)
: ShapeGeneralizedTree(criterion, numPartitions, outerParams, innerParams),
numClasses_(std::move(numClasses)), cdParams_(cdParams),
random_state_(random_state), rng_(),
featureBagging_(featureBagging
? std::move(featureBagging)
: FeatureBaggingPickFn(pickAllFeatureIndices)),
+ pairwiseCandidates_(pairwiseCandidates), pairwisePenalty_(pairwisePenalty),
outerTreeBuilder_(outerParams_.minLeafSize, outerParams_.minGainSplit,
outerParams_.maxDepth, outerParams_.maxLeafNodes) {
if (criterion != LearningCriterion::Entropy &&
@@ -72,6 +75,14 @@ ClassificationShapeGeneralizedTree::ClassificationShapeGeneralizedTree(
if (numPartitions < 2)
throw std::invalid_argument(
"ClassificationShapeGeneralizedTree: numPartitions must be >= 2");
+ if (!std::isfinite(pairwisePenalty_) || pairwisePenalty_ < 0.0)
+ throw std::invalid_argument("pairwise_penalty must be finite and non-negative");
+}
+
+bool ClassificationShapeGeneralizedTree::hasPairNodes() const {
+ return std::any_of(nodes_.begin(), nodes_.end(), [](const ShapeFunctionNode &node) {
+ return !node.isLeaf && node.logicalFeatureIndices.size() == 2;
+ });
}
void ClassificationShapeGeneralizedTree::resolveOutputLayout(size_t nOutputs) {
@@ -214,6 +225,23 @@ void ClassificationShapeGeneralizedTree::fit(
const size_t xSubCols = static_cast(Xsub.n_cols);
ShapeBestBranchingState best{};
+ const arma::Row wsub =
+ subSampleWeights(fitSampleWeights_, subIdx);
+
+ struct UnivariateProxy {
+ size_t logicalIndex;
+ size_t numPartitions;
+ double childImpurity;
+ std::vector partitions;
+ };
+ std::vector univariateProxies;
+ std::vector retainedPairCandidates;
+
+ const auto addNoSplitProxy = [&univariateProxies, xSubCols,
+ parentImp](size_t logicalIdx) {
+ univariateProxies.push_back(
+ {logicalIdx, 1, parentImp, std::vector(xSubCols, 0)});
+ };
const auto applyTaskFields =
[](ShapeBestBranchingState &state,
@@ -228,24 +256,39 @@ void ClassificationShapeGeneralizedTree::fit(
const size_t logicalIdx = featureSubset[fi];
const FeatureInfo &feature = features_[logicalIdx];
- const arma::Row wsub =
- subSampleWeights(fitSampleWeights_, subIdx);
-
auto disc = makeClassificationDiscretizer(criterion_, feature);
trainClassificationDiscretizer(
*disc, feature, Xsub, ysub, classesPerOutput_,
innerParams_.minLeafSize, innerParams_.minGainSplit,
innerParams_.maxDepth, innerParams_.maxLeafNodes, wsub);
- if (disc->numLeaves() < 2)
+ if (disc->numLeaves() < 2) {
+ if (pairwiseCandidates_ > 0)
+ addNoSplitProxy(logicalIdx);
continue;
+ }
const ShapeBranchAssignmentSearchResult featureBest =
searchShapeBranchAssignmentFromDiscretizer(
*disc, criterion_, parentImp, numPartitions_, outerParams_,
cdParams_, outerTreeBuilder_.eps, rng_,
/*useKMeansSeed=*/true, classesPerOutput_, nOutputs_);
- if (!featureBest.found)
+ if (!featureBest.found) {
+ if (pairwiseCandidates_ > 0)
+ addNoSplitProxy(logicalIdx);
continue;
+ }
+
+ if (pairwiseCandidates_ > 0) {
+ std::vector partitions(xSubCols, 0);
+ const auto &perBin = disc->inSampleDiscretizations();
+ for (size_t bin = 0; bin < perBin.size(); ++bin)
+ for (size_t sample : perBin[bin])
+ partitions[sample] = featureBest.assignments[bin];
+ univariateProxies.push_back(
+ {logicalIdx, featureBest.chosenK,
+ parentImp - featureBest.impurityDecrease,
+ std::move(partitions)});
+ }
featureHasBetterShapeBranching(
featureBest, best, logicalIdx, xSubCols, feature.indices,
@@ -254,6 +297,93 @@ void ClassificationShapeGeneralizedTree::fit(
outerTreeBuilder_.eps, applyTaskFields);
}
+ if (pairwiseCandidates_ > 0 && univariateProxies.size() >= 2) {
+ std::sort(univariateProxies.begin(), univariateProxies.end(),
+ [](const UnivariateProxy &a, const UnivariateProxy &b) {
+ return a.logicalIndex < b.logicalIndex;
+ });
+ struct PairProxy {
+ double score;
+ size_t first;
+ size_t second;
+ };
+ std::vector retained;
+ const auto proxyLess = [](const PairProxy &a, const PairProxy &b) {
+ if (a.score != b.score)
+ return a.score < b.score;
+ if (a.first != b.first)
+ return a.first < b.first;
+ return a.second < b.second;
+ };
+ const double totalWeight = arma::accu(wsub);
+ for (size_t i = 0; i + 1 < univariateProxies.size(); ++i) {
+ for (size_t j = i + 1; j < univariateProxies.size(); ++j) {
+ const size_t numCells = univariateProxies[i].numPartitions *
+ univariateProxies[j].numPartitions;
+ auto crossed = std::vector>>(
+ numCells, makeEmptyHistogram());
+ std::vector crossedWeights(numCells, 0.0);
+ for (size_t sample = 0; sample < xSubCols; ++sample) {
+ const size_t cell =
+ univariateProxies[i].partitions[sample] *
+ univariateProxies[j].numPartitions +
+ univariateProxies[j].partitions[sample];
+ const double w = wsub(sample);
+ crossedWeights[cell] += w;
+ for (size_t o = 0; o < nOutputs_; ++o)
+ crossed[cell][o][ysub(o, sample)] += w;
+ }
+ double crossedImpurity = 0.0;
+ for (size_t cell = 0; cell < crossed.size(); ++cell)
+ if (totalWeight > 0.0)
+ crossedImpurity += crossedWeights[cell] / totalWeight *
+ impurityForClassCounts(crossed[cell]);
+ retained.push_back(
+ {crossedImpurity -
+ std::min(univariateProxies[i].childImpurity,
+ univariateProxies[j].childImpurity),
+ univariateProxies[i].logicalIndex,
+ univariateProxies[j].logicalIndex});
+ std::sort(retained.begin(), retained.end(), proxyLess);
+ if (retained.size() > pairwiseCandidates_)
+ retained.resize(pairwiseCandidates_);
+ }
+ }
+
+ retainedPairCandidates.reserve(retained.size());
+ for (const PairProxy &pair : retained)
+ retainedPairCandidates.push_back(
+ {{pair.first, pair.second},
+ {features_[pair.first], features_[pair.second]}});
+
+ for (const PairProxy &pair : retained) {
+ const FeatureInfo &first = features_[pair.first];
+ const FeatureInfo &second = features_[pair.second];
+ arma::uvec rawFeatures = arma::join_cols(first.indices, second.indices);
+ auto pairDisc = std::make_unique(
+ criterion_, first, second);
+ pairDisc->Train(
+ Xsub, rawFeatures, ysub, classesPerOutput_,
+ innerParams_.minLeafSize, innerParams_.minGainSplit,
+ innerParams_.maxDepth, innerParams_.maxLeafNodes, wsub);
+ if (pairDisc->numLeaves() < 2)
+ continue;
+ ShapeBranchAssignmentSearchResult pairBest =
+ searchShapeBranchAssignmentFromDiscretizer(
+ *pairDisc, criterion_, parentImp, numPartitions_,
+ outerParams_, cdParams_, outerTreeBuilder_.eps, rng_,
+ /*useKMeansSeed=*/true, classesPerOutput_, nOutputs_,
+ nullptr, nullptr, 0, /*hasNanRoutingBin=*/false);
+ pairBest.bestFeatureScore += pairwisePenalty_;
+ if (featureHasBetterShapeBranching(
+ pairBest, best, pair.first, xSubCols, rawFeatures,
+ std::unique_ptr>>(
+ std::move(pairDisc)),
+ outerTreeBuilder_.eps, applyTaskFields))
+ best.logicalFeatureIndices = {pair.first, pair.second};
+ }
+ }
+
if (!std::isfinite(best.penalizedChildScore) ||
best.penalizedChildScore >= std::numeric_limits::infinity() ||
best.branching.impurityDecrease <= outerTreeBuilder_.eps) {
@@ -263,6 +393,8 @@ void ClassificationShapeGeneralizedTree::fit(
node.isLeaf = false;
node.splitFeatureIndex = best.branching.featureIndex;
+ node.logicalFeatureIndices = best.logicalFeatureIndices;
+ node.retainedPairCandidates = std::move(retainedPairCandidates);
node.routingFeatures.assign(best.routingColumnIndices.begin(),
best.routingColumnIndices.end());
node.innerDiscretizer = best.winningDiscretizer;
@@ -313,8 +445,16 @@ void ClassificationShapeGeneralizedTree::fit(
nodes_[0], findBestSplit, makeChildren,
[this](ShapeFunctionNode &parent,
std::vector &children) {
- sumOfNodeImportancesByFeature_(parent.splitFeatureIndex) +=
- parent.informationGain;
+ const auto &logical = parent.logicalFeatureIndices;
+ if (logical.size() == 2) {
+ sumOfNodeImportancesByFeature_(logical[0]) +=
+ parent.informationGain / 2.0;
+ sumOfNodeImportancesByFeature_(logical[1]) +=
+ parent.informationGain / 2.0;
+ } else {
+ sumOfNodeImportancesByFeature_(parent.splitFeatureIndex) +=
+ parent.informationGain;
+ }
totalNodeImportanceSum_ += parent.informationGain;
const size_t pid = parent.nodeIndex;
nodes_[pid] = std::move(parent);
diff --git a/cpp/src/Estimators/ClassificationShapeGeneralizedTree.h b/cpp/src/Estimators/ClassificationShapeGeneralizedTree.h
index dfe7950..57e8648 100644
--- a/cpp/src/Estimators/ClassificationShapeGeneralizedTree.h
+++ b/cpp/src/Estimators/ClassificationShapeGeneralizedTree.h
@@ -86,7 +86,8 @@ class ClassificationShapeGeneralizedTree : public ShapeGeneralizedTree {
size_t numPartitions, TreeBuildingParams outerParams = {},
TreeBuildingParams innerParams = {},
CoordinateDescentParams cdParams = {}, uint64_t random_state = 42,
- FeatureBaggingPickFn featureBagging = {});
+ FeatureBaggingPickFn featureBagging = {}, size_t pairwiseCandidates = 0,
+ double pairwisePenalty = 0.0);
/** Convenience overload: single-output / shared class count ``{numClasses}``. */
ClassificationShapeGeneralizedTree(
@@ -94,11 +95,12 @@ class ClassificationShapeGeneralizedTree : public ShapeGeneralizedTree {
TreeBuildingParams outerParams = {},
TreeBuildingParams innerParams = {},
CoordinateDescentParams cdParams = {}, uint64_t random_state = 42,
- FeatureBaggingPickFn featureBagging = {})
+ FeatureBaggingPickFn featureBagging = {}, size_t pairwiseCandidates = 0,
+ double pairwisePenalty = 0.0)
: ClassificationShapeGeneralizedTree(
criterion, std::vector{numClasses}, numPartitions,
outerParams, innerParams, cdParams, random_state,
- std::move(featureBagging)) {}
+ std::move(featureBagging), pairwiseCandidates, pairwisePenalty) {}
~ClassificationShapeGeneralizedTree() = default;
@@ -151,6 +153,8 @@ class ClassificationShapeGeneralizedTree : public ShapeGeneralizedTree {
/** Number of outputs the tree was fitted on (>= 1). */
size_t nOutputs() const { return nOutputs_; }
+ bool hasPairNodes() const;
+
/** Fit-resolved per-output class counts (empty before ``fit``). */
const std::vector &classesPerOutput() const {
return classesPerOutput_;
@@ -167,6 +171,8 @@ class ClassificationShapeGeneralizedTree : public ShapeGeneralizedTree {
uint64_t random_state_;
std::mt19937_64 rng_;
FeatureBaggingPickFn featureBagging_;
+ size_t pairwiseCandidates_ = 0;
+ double pairwisePenalty_ = 0.0;
std::vector features_;
/** Outer routing expansion; `fit` passes split logic via buildTree callbacks. */
diff --git a/cpp/src/Estimators/RegressionShapeGeneralizedTree.cpp b/cpp/src/Estimators/RegressionShapeGeneralizedTree.cpp
index 8d6e660..a2aad0c 100644
--- a/cpp/src/Estimators/RegressionShapeGeneralizedTree.cpp
+++ b/cpp/src/Estimators/RegressionShapeGeneralizedTree.cpp
@@ -13,6 +13,7 @@
#include "Estimators/RegressionShapeGeneralizedTree.h"
#include "Criterion.h"
+#include "Discretizers/pair/PairRegressionDiscretizer.h"
#include "Discretizers/factories/DiscretizerFactories.h"
#include "Discretizers/RegressionDiscretizer.h"
#include "Estimators/ShapeFunctions/ShapeFunctionSplitSearch.h"
@@ -65,19 +66,77 @@ PartitionMoments aggregatePartitionFromBins(
return out;
}
+double crossedRegressionImpurity(
+ const std::vector &firstPartitions, size_t firstK,
+ const std::vector &secondPartitions, size_t secondK,
+ const arma::Mat &y, const arma::Row &weights,
+ LearningCriterion criterion, size_t nOutputs) {
+ const size_t numCells = firstK * secondK;
+ std::vector cellWeights(numCells, 0.0);
+ double totalWeight = 0.0;
+ for (size_t sample = 0; sample < firstPartitions.size(); ++sample) {
+ const size_t cell = firstPartitions[sample] * secondK + secondPartitions[sample];
+ const double w = weights(sample);
+ cellWeights[cell] += w;
+ totalWeight += w;
+ }
+ if (totalWeight <= 0.0)
+ return 0.0;
+
+ if (criterion == LearningCriterion::SquaredError) {
+ std::vector>> stats(
+ numCells, std::vector>(nOutputs,
+ std::vector(2, 0.0)));
+ for (size_t sample = 0; sample < firstPartitions.size(); ++sample) {
+ const size_t cell = firstPartitions[sample] * secondK + secondPartitions[sample];
+ const double w = weights(sample);
+ for (size_t o = 0; o < nOutputs; ++o) {
+ const double value = y(o, sample);
+ stats[cell][o][0] += w * value;
+ stats[cell][o][1] += w * value * value;
+ }
+ }
+ double result = 0.0;
+ for (size_t cell = 0; cell < numCells; ++cell)
+ result += cellWeights[cell] / totalWeight *
+ Criterion::squaredError(stats[cell], cellWeights[cell]);
+ return result;
+ }
+
+ std::vector>> cellYs(
+ numCells, std::vector>(nOutputs));
+ std::vector> cellWs(numCells);
+ for (size_t sample = 0; sample < firstPartitions.size(); ++sample) {
+ const size_t cell = firstPartitions[sample] * secondK + secondPartitions[sample];
+ cellWs[cell].push_back(weights(sample));
+ for (size_t o = 0; o < nOutputs; ++o)
+ cellYs[cell][o].push_back(y(o, sample));
+ }
+ double result = 0.0;
+ for (size_t cell = 0; cell < numCells; ++cell) {
+ double cellImpurity = 0.0;
+ for (size_t o = 0; o < nOutputs; ++o)
+ cellImpurity += Criterion::absoluteError(cellYs[cell][o], cellWs[cell]).mae;
+ result += cellWeights[cell] / totalWeight * cellImpurity;
+ }
+ return result;
+}
+
} // namespace
RegressionShapeGeneralizedTree::RegressionShapeGeneralizedTree(
LearningCriterion criterion, size_t numPartitions,
TreeBuildingParams outerParams, TreeBuildingParams innerParams,
CoordinateDescentParams cdParams, uint64_t random_state,
- FeatureBaggingPickFn featureBagging)
+ FeatureBaggingPickFn featureBagging, size_t pairwiseCandidates,
+ double pairwisePenalty)
: ShapeGeneralizedTree(criterion, numPartitions, outerParams, innerParams),
cdParams_(cdParams),
random_state_(random_state), rng_(),
featureBagging_(featureBagging
? std::move(featureBagging)
: FeatureBaggingPickFn(pickAllFeatureIndices)),
+ pairwiseCandidates_(pairwiseCandidates), pairwisePenalty_(pairwisePenalty),
outerTreeBuilder_(outerParams_.minLeafSize, outerParams_.minGainSplit,
outerParams_.maxDepth, outerParams_.maxLeafNodes) {
if (criterion != LearningCriterion::SquaredError &&
@@ -88,6 +147,14 @@ RegressionShapeGeneralizedTree::RegressionShapeGeneralizedTree(
if (numPartitions < 2)
throw std::invalid_argument(
"RegressionShapeGeneralizedTree: numPartitions must be >= 2");
+ if (!std::isfinite(pairwisePenalty_) || pairwisePenalty_ < 0.0)
+ throw std::invalid_argument("pairwise_penalty must be finite and non-negative");
+}
+
+bool RegressionShapeGeneralizedTree::hasPairNodes() const {
+ return std::any_of(nodes_.begin(), nodes_.end(), [](const ShapeFunctionNode &node) {
+ return !node.isLeaf && node.logicalFeatureIndices.size() == 2;
+ });
}
@@ -279,6 +346,22 @@ void RegressionShapeGeneralizedTree::fit(
const size_t xSubCols = static_cast(Xsub.n_cols);
ShapeBestBranchingState best{};
+ const arma::Row wsub =
+ subSampleWeights(fitSampleWeights_, subIdx);
+ struct UnivariateProxy {
+ size_t logicalIndex;
+ size_t numPartitions;
+ double childImpurity;
+ std::vector partitions;
+ };
+ std::vector univariateProxies;
+ std::vector retainedPairCandidates;
+
+ const auto addNoSplitProxy = [&univariateProxies, xSubCols,
+ parentImp](size_t logicalIdx) {
+ univariateProxies.push_back(
+ {logicalIdx, 1, parentImp, std::vector(xSubCols, 0)});
+ };
const auto applyTaskFields =
[this](ShapeBestBranchingState &state,
@@ -295,16 +378,16 @@ void RegressionShapeGeneralizedTree::fit(
const size_t logicalIdx = featureSubset[fi];
const FeatureInfo &feature = features_[logicalIdx];
- const arma::Row wsub =
- subSampleWeights(fitSampleWeights_, subIdx);
-
auto disc = makeRegressionDiscretizer(criterion_, feature);
trainRegressionDiscretizer(
*disc, feature, Xsub, ysub, innerParams_.minLeafSize,
innerParams_.minGainSplit, innerParams_.maxDepth,
innerParams_.maxLeafNodes, wsub);
- if (disc->numLeaves() < 2)
+ if (disc->numLeaves() < 2) {
+ if (pairwiseCandidates_ > 0)
+ addNoSplitProxy(logicalIdx);
continue;
+ }
const ShapeBranchAssignmentSearchResult featureBest =
searchShapeBranchAssignmentFromDiscretizer(
@@ -316,8 +399,23 @@ void RegressionShapeGeneralizedTree::fit(
criterion_ == LearningCriterion::AbsoluteError ? &wsub
: nullptr,
xSubCols);
- if (!featureBest.found)
+ if (!featureBest.found) {
+ if (pairwiseCandidates_ > 0)
+ addNoSplitProxy(logicalIdx);
continue;
+ }
+
+ if (pairwiseCandidates_ > 0) {
+ std::vector partitions(xSubCols, 0);
+ const auto &perBin = disc->inSampleDiscretizations();
+ for (size_t bin = 0; bin < perBin.size(); ++bin)
+ for (size_t sample : perBin[bin])
+ partitions[sample] = featureBest.assignments[bin];
+ univariateProxies.push_back(
+ {logicalIdx, featureBest.chosenK,
+ parentImp - featureBest.impurityDecrease,
+ std::move(partitions)});
+ }
featureHasBetterShapeBranching(
featureBest, best, logicalIdx, xSubCols, feature.indices,
@@ -326,6 +424,77 @@ void RegressionShapeGeneralizedTree::fit(
outerTreeBuilder_.eps, applyTaskFields);
}
+ if (pairwiseCandidates_ > 0 && univariateProxies.size() >= 2) {
+ std::sort(univariateProxies.begin(), univariateProxies.end(),
+ [](const UnivariateProxy &a, const UnivariateProxy &b) {
+ return a.logicalIndex < b.logicalIndex;
+ });
+ struct PairProxy {
+ double score;
+ size_t first;
+ size_t second;
+ };
+ const auto proxyLess = [](const PairProxy &a, const PairProxy &b) {
+ if (a.score != b.score)
+ return a.score < b.score;
+ if (a.first != b.first)
+ return a.first < b.first;
+ return a.second < b.second;
+ };
+ std::vector retained;
+ for (size_t i = 0; i + 1 < univariateProxies.size(); ++i) {
+ for (size_t j = i + 1; j < univariateProxies.size(); ++j) {
+ const double crossed = crossedRegressionImpurity(
+ univariateProxies[i].partitions, univariateProxies[i].numPartitions,
+ univariateProxies[j].partitions, univariateProxies[j].numPartitions,
+ ysub, wsub, criterion_, nOutputs_);
+ retained.push_back(
+ {crossed - std::min(univariateProxies[i].childImpurity,
+ univariateProxies[j].childImpurity),
+ univariateProxies[i].logicalIndex,
+ univariateProxies[j].logicalIndex});
+ std::sort(retained.begin(), retained.end(), proxyLess);
+ if (retained.size() > pairwiseCandidates_)
+ retained.resize(pairwiseCandidates_);
+ }
+ }
+
+ retainedPairCandidates.reserve(retained.size());
+ for (const PairProxy &pair : retained)
+ retainedPairCandidates.push_back(
+ {{pair.first, pair.second},
+ {features_[pair.first], features_[pair.second]}});
+
+ for (const PairProxy &pair : retained) {
+ const FeatureInfo &first = features_[pair.first];
+ const FeatureInfo &second = features_[pair.second];
+ arma::uvec rawFeatures = arma::join_cols(first.indices, second.indices);
+ auto pairDisc = std::make_unique(
+ criterion_, first, second);
+ pairDisc->Train(
+ Xsub, rawFeatures, ysub, innerParams_.minLeafSize,
+ innerParams_.minGainSplit, innerParams_.maxDepth,
+ innerParams_.maxLeafNodes, wsub);
+ if (pairDisc->numLeaves() < 2)
+ continue;
+ ShapeBranchAssignmentSearchResult pairBest =
+ searchShapeBranchAssignmentFromDiscretizer(
+ *pairDisc, criterion_, parentImp, numPartitions_, outerParams_,
+ cdParams_, outerTreeBuilder_.eps, rng_,
+ /*useKMeansSeed=*/false, /*classesPerOutput=*/{}, nOutputs_,
+ criterion_ == LearningCriterion::AbsoluteError ? &ysub : nullptr,
+ criterion_ == LearningCriterion::AbsoluteError ? &wsub : nullptr,
+ xSubCols, /*hasNanRoutingBin=*/false);
+ pairBest.bestFeatureScore += pairwisePenalty_;
+ if (featureHasBetterShapeBranching(
+ pairBest, best, pair.first, xSubCols, rawFeatures,
+ std::unique_ptr>>(
+ std::move(pairDisc)),
+ outerTreeBuilder_.eps, applyTaskFields))
+ best.logicalFeatureIndices = {pair.first, pair.second};
+ }
+ }
+
if (!std::isfinite(best.penalizedChildScore) ||
best.penalizedChildScore >= std::numeric_limits::infinity() ||
best.branching.impurityDecrease <= outerTreeBuilder_.eps) {
@@ -335,6 +504,8 @@ void RegressionShapeGeneralizedTree::fit(
node.isLeaf = false;
node.splitFeatureIndex = best.branching.featureIndex;
+ node.logicalFeatureIndices = best.logicalFeatureIndices;
+ node.retainedPairCandidates = std::move(retainedPairCandidates);
node.routingFeatures.assign(best.routingColumnIndices.begin(),
best.routingColumnIndices.end());
node.innerDiscretizer = best.winningDiscretizer;
@@ -423,8 +594,15 @@ void RegressionShapeGeneralizedTree::fit(
nodes_[0], findBestSplit, makeChildren,
[this](ShapeFunctionNode &parent,
std::vector &children) {
- sumOfNodeImportancesByFeature_(parent.splitFeatureIndex) +=
- parent.informationGain;
+ if (parent.logicalFeatureIndices.size() == 2) {
+ sumOfNodeImportancesByFeature_(parent.logicalFeatureIndices[0]) +=
+ parent.informationGain / 2.0;
+ sumOfNodeImportancesByFeature_(parent.logicalFeatureIndices[1]) +=
+ parent.informationGain / 2.0;
+ } else {
+ sumOfNodeImportancesByFeature_(parent.splitFeatureIndex) +=
+ parent.informationGain;
+ }
totalNodeImportanceSum_ += parent.informationGain;
const size_t pid = parent.nodeIndex;
nodes_[pid] = parent;
@@ -508,5 +686,3 @@ RegressionShapeGeneralizedTree::predict(const arma::fmat &X) const {
}
return yhat;
}
-
-
diff --git a/cpp/src/Estimators/RegressionShapeGeneralizedTree.h b/cpp/src/Estimators/RegressionShapeGeneralizedTree.h
index 28c8369..3ecf886 100644
--- a/cpp/src/Estimators/RegressionShapeGeneralizedTree.h
+++ b/cpp/src/Estimators/RegressionShapeGeneralizedTree.h
@@ -80,7 +80,8 @@ class RegressionShapeGeneralizedTree : public ShapeGeneralizedTree {
TreeBuildingParams outerParams = {},
TreeBuildingParams innerParams = {},
CoordinateDescentParams cdParams = {}, uint64_t random_state = 42,
- FeatureBaggingPickFn featureBagging = {});
+ FeatureBaggingPickFn featureBagging = {}, size_t pairwiseCandidates = 0,
+ double pairwisePenalty = 0.0);
~RegressionShapeGeneralizedTree() = default;
@@ -114,6 +115,8 @@ class RegressionShapeGeneralizedTree : public ShapeGeneralizedTree {
/** Number of outputs the tree was fitted on (>= 1). */
size_t nOutputs() const { return nOutputs_; }
+ bool hasPairNodes() const;
+
/**
* Per outer-tree node index: for squared error, concatenated
* ``[Σw·y0, Σw·y0², Σw·y1, Σw·y1², ...]`` (length ``2 * nOutputs``) at leaves
@@ -142,6 +145,8 @@ class RegressionShapeGeneralizedTree : public ShapeGeneralizedTree {
uint64_t random_state_;
std::mt19937_64 rng_;
FeatureBaggingPickFn featureBagging_;
+ size_t pairwiseCandidates_ = 0;
+ double pairwisePenalty_ = 0.0;
std::vector features_;
/** Outer routing expansion; `fit` passes split logic via buildTree callbacks. */
diff --git a/cpp/src/Estimators/ShapeFunctions/ShapeFunctionNode.h b/cpp/src/Estimators/ShapeFunctions/ShapeFunctionNode.h
index bfef499..e7ec081 100644
--- a/cpp/src/Estimators/ShapeFunctions/ShapeFunctionNode.h
+++ b/cpp/src/Estimators/ShapeFunctions/ShapeFunctionNode.h
@@ -8,8 +8,10 @@
*/
#include "Discretizers/InnerDiscretizerBase.h"
+#include "Domain/FeatureInfo.h"
#include
+#include
#include
#include
#include
@@ -18,6 +20,11 @@
class ShapeGeneralizedTree;
+struct RetainedPairCandidate {
+ std::array logicalFeatureIndices{};
+ std::array features;
+};
+
/**
* Outer-tree node: routing rule when internal, plus training sample indices
* during fit.
@@ -48,6 +55,12 @@ class ShapeFunctionNode {
*/
size_t splitFeatureIndex = 0;
+ /** Logical feature indices used by this router (one normally, two for a pair). */
+ std::vector logicalFeatureIndices;
+
+ /** Top-P pairs from initial screening; TAO may refit only these pairs. */
+ std::vector retainedPairCandidates;
+
/** Row indices into X used for routing; undefined if isLeaf. */
std::vector routingFeatures;
/** Maps each inner discretizer bin (including NaN) to a child partition. */
diff --git a/cpp/src/Estimators/ShapeFunctions/ShapeFunctionSplitSearch.cpp b/cpp/src/Estimators/ShapeFunctions/ShapeFunctionSplitSearch.cpp
index d17dc37..3b25837 100644
--- a/cpp/src/Estimators/ShapeFunctions/ShapeFunctionSplitSearch.cpp
+++ b/cpp/src/Estimators/ShapeFunctions/ShapeFunctionSplitSearch.cpp
@@ -28,14 +28,16 @@ void refineShapeBranchAssignmentNested(
std::vector>> &stats,
std::vector &leafWeights,
const std::vector &leafSampleCounts,
- const std::vector &classesPerOutput, size_t nOutputs) {
+ const std::vector &classesPerOutput, size_t nOutputs,
+ bool hasNanRoutingBin) {
if (k >= numRoutingBins ||
criterion == LearningCriterion::AbsoluteError)
return;
const std::vector snapshot = branchObj->assignments;
const double objBeforeCd = branchObj->objective();
- coordinateDescent(k, *branchObj, rng, cdParams.maxIters, cdParams.patience);
+ coordinateDescent(k, *branchObj, rng, cdParams.maxIters, cdParams.patience,
+ hasNanRoutingBin);
const double objAfterCd = branchObj->objective();
if (std::isfinite(objAfterCd) &&
objAfterCd <= objBeforeCd + kShapeFunctionCdImprovementEps)
@@ -53,13 +55,14 @@ void refineShapeBranchAssignmentAbsoluteError(
std::vector>> &maeLeafYs,
std::vector> &maeLeafWs,
std::vector &leafWeights,
- const std::vector &leafSampleCounts) {
+ const std::vector &leafSampleCounts, bool hasNanRoutingBin) {
if (k >= numRoutingBins || !mae_branch_config::coordinateDescentEnabled())
return;
const std::vector snapshot = branchObj->assignments;
const double objBeforeCd = branchObj->objective();
- coordinateDescent(k, *branchObj, rng, cdParams.maxIters, cdParams.patience);
+ coordinateDescent(k, *branchObj, rng, cdParams.maxIters, cdParams.patience,
+ hasNanRoutingBin);
const double objAfterCd = branchObj->objective();
if (std::isfinite(objAfterCd) &&
objAfterCd <= objBeforeCd + kShapeFunctionCdImprovementEps)
@@ -117,7 +120,7 @@ ShapeBranchAssignmentSearchResult searchShapeBranchAssignmentFromDiscretizer(
std::mt19937_64 &rng, bool useKMeansSeed,
const std::vector &classesPerOutput, size_t nOutputs,
const arma::Mat *ysub, const arma::Row *wsub,
- size_t xSubCols) {
+ size_t xSubCols, bool hasNanRoutingBin) {
auto &stats = disc.leafStats();
auto &sizes = disc.leafNumSamples();
auto &leafWeights = disc.leafNodeWeights();
@@ -176,14 +179,14 @@ ShapeBranchAssignmentSearchResult searchShapeBranchAssignmentFromDiscretizer(
maeLeafYs, maeLeafWs);
refineShapeBranchAssignmentAbsoluteError(
branchObj, k, numRoutingBins, cdParams, rng, maeLeafYsStorage,
- maeLeafWsStorage, leafWeights, sizes);
+ maeLeafWsStorage, leafWeights, sizes, hasNanRoutingBin);
} else {
branchObj = makeBranchAssignment(criterion, trialAssignments, k, stats,
leafWeights, sizes, classesPerOutput,
nOutputs);
refineShapeBranchAssignmentNested(
branchObj, k, numRoutingBins, criterion, cdParams, rng, stats,
- leafWeights, sizes, classesPerOutput, nOutputs);
+ leafWeights, sizes, classesPerOutput, nOutputs, hasNanRoutingBin);
}
if (!branchObj->partitionCountsMeetMinLeaf(outerParams.minLeafSize))
@@ -199,7 +202,10 @@ ShapeBranchAssignmentSearchResult searchShapeBranchAssignmentFromDiscretizer(
if (score < result.bestFeatureScore - scoreEpsilon) {
result.bestFeatureScore = score;
result.chosenK = k;
- result.assignments = trialAssignments;
+ // Coordinate descent updates the live branch object, not the seed vector.
+ // Persist that refined routing so the objective and eventual child buckets
+ // describe the same split.
+ result.assignments = branchObj->assignments;
result.partitionSampleCounts = branchObj->partitionSampleCounts();
if (const auto *leafAgg =
dynamic_castleafNodeWeights().end());
applyTaskFields(best, search, disc->leafStats());
best.routingColumnIndices = routingColumnIndices;
+ best.logicalFeatureIndices = {featureIndex};
best.winningDiscretizer =
std::shared_ptr(std::move(disc));
return true;
diff --git a/cpp/src/Estimators/ShapeFunctions/ShapeFunctionSplitSearch.h b/cpp/src/Estimators/ShapeFunctions/ShapeFunctionSplitSearch.h
index 5838371..25d5a16 100644
--- a/cpp/src/Estimators/ShapeFunctions/ShapeFunctionSplitSearch.h
+++ b/cpp/src/Estimators/ShapeFunctions/ShapeFunctionSplitSearch.h
@@ -46,6 +46,7 @@ struct ShapeBestBranchingState {
std::shared_ptr winningDiscretizer;
/** Column indices into ``X`` used for routing at inference. */
arma::uvec routingColumnIndices;
+ std::vector logicalFeatureIndices;
};
void markShapeFunctionNodeAsLeaf(ShapeFunctionNode &node);
@@ -87,7 +88,7 @@ ShapeBranchAssignmentSearchResult searchShapeBranchAssignmentFromDiscretizer(
std::mt19937_64 &rng, bool useKMeansSeed = false,
const std::vector &classesPerOutput = {}, size_t nOutputs = 1,
const arma::Mat *ysub = nullptr, const arma::Row *wsub = nullptr,
- size_t xSubCols = 0);
+ size_t xSubCols = 0, bool hasNanRoutingBin = true);
std::vector>
routeSamplesToPartitions(const ShapeFunctionNode &parent, const arma::fmat &X);
diff --git a/cpp/src/algorithms/CoordinateDescent.h b/cpp/src/algorithms/CoordinateDescent.h
index 6c03a49..7eac238 100644
--- a/cpp/src/algorithms/CoordinateDescent.h
+++ b/cpp/src/algorithms/CoordinateDescent.h
@@ -30,22 +30,22 @@
inline double coordinateDescent(size_t numPartitions,
BranchAssignment &assignmentObjective,
std::mt19937_64 &rng, size_t maxIters = 10,
- size_t patience = 5) {
+ size_t patience = 5,
+ bool hasNanRoutingBin = true) {
size_t numBins = assignmentObjective.assignments.size();
if (numBins <= 1) return assignmentObjective.objective();
- const size_t nanBinIndex = numBins - 1;
+ const size_t finiteBinCount = hasNanRoutingBin ? numBins - 1 : numBins;
- // 1. Omit the NaN bin from the objective state during standard coordinate descent
- assignmentObjective.removeLeaf(nanBinIndex);
+ if (hasNanRoutingBin)
+ assignmentObjective.removeLeaf(numBins - 1);
size_t consecutiveTrialsWithoutImprovement = 0;
for (size_t i = 0; i < maxIters; ++i) {
bool improved = false;
- // 2. Shuffle and optimize ONLY the finite numeric bins (0 to numBins - 2)
- std::vector permutation(nanBinIndex);
+ std::vector permutation(finiteBinCount);
std::iota(permutation.begin(), permutation.end(), size_t{0});
std::shuffle(permutation.begin(), permutation.end(), rng);
@@ -84,7 +84,11 @@ inline double coordinateDescent(size_t numPartitions,
}
- // 3. Factor the NaN bin back in by greedily finding its optimal partition
+ if (!hasNanRoutingBin)
+ return assignmentObjective.objective();
+
+ // Factor the NaN bin back in by greedily finding its optimal partition.
+ const size_t nanBinIndex = numBins - 1;
size_t bestNanPartition = missing_values::partition_with_max_count_min_index_tie(assignmentObjective.partitionSampleCounts()); // Fallback
assignmentObjective.addLeaf(
nanBinIndex,
@@ -110,4 +114,4 @@ inline double coordinateDescent(size_t numPartitions,
assignmentObjective.addLeaf(nanBinIndex, bestNanPartition);
return assignmentObjective.objective();
-}
\ No newline at end of file
+}
diff --git a/cpp/src/algorithms/TAO/ClassificationTaoAdapter.cpp b/cpp/src/algorithms/TAO/ClassificationTaoAdapter.cpp
index 92ca451..2384df6 100644
--- a/cpp/src/algorithms/TAO/ClassificationTaoAdapter.cpp
+++ b/cpp/src/algorithms/TAO/ClassificationTaoAdapter.cpp
@@ -178,4 +178,33 @@ void ClassificationTaoAdapter::recomputeLeafStats(
classesPerOutput_, nOutputs_);
}
+void ClassificationTaoAdapter::refreshNodeBinMetadata(
+ ShapeFunctionNode &node, const std::vector &samples) {
+ if (!node.innerDiscretizer)
+ throw std::runtime_error(
+ "ClassificationTaoAdapter: accepted router has no discretizer");
+ arma::Row bins;
+ node.innerDiscretizer->transform(X_, bins);
+ const size_t nBins = node.binToPartition.size();
+ node.binSampleCounts.assign(nBins, 0);
+ node.splitBinWeights.assign(nBins, 0.0);
+ node.splitClassCounts.assign(nBins, std::vector>(nOutputs_));
+ for (size_t b = 0; b < nBins; ++b)
+ for (size_t o = 0; o < nOutputs_; ++o)
+ node.splitClassCounts[b][o].assign(classesPerOutput_[o], 0.0);
+ for (arma::uword col : samples) {
+ const size_t bin = bins(col);
+ if (bin >= nBins)
+ throw std::runtime_error(
+ "ClassificationTaoAdapter: accepted router bin out of range");
+ ++node.binSampleCounts[bin];
+ const double weight = static_cast(w_(col));
+ node.splitBinWeights[bin] += weight;
+ for (size_t o = 0; o < nOutputs_; ++o)
+ node.splitClassCounts[bin][o][y_(static_cast(o), col)] +=
+ weight;
+ }
+ node.splitLeafStats.clear();
+}
+
} // namespace tao
diff --git a/cpp/src/algorithms/TAO/ClassificationTaoAdapter.h b/cpp/src/algorithms/TAO/ClassificationTaoAdapter.h
index b2fe80a..9443fb5 100644
--- a/cpp/src/algorithms/TAO/ClassificationTaoAdapter.h
+++ b/cpp/src/algorithms/TAO/ClassificationTaoAdapter.h
@@ -51,6 +51,10 @@ class ClassificationTaoAdapter final : public ShapeGeneralizedTaoAdapter {
void recomputeLeafStats(
const std::vector> &nodeSamples) override;
+ void refreshNodeBinMetadata(
+ ShapeFunctionNode &node,
+ const std::vector &samples) override;
+
private:
/**
* Per-child correctness rewards for one sample.
diff --git a/cpp/src/algorithms/TAO/RegressionTaoAdapter.cpp b/cpp/src/algorithms/TAO/RegressionTaoAdapter.cpp
index 2ce16b1..5e1e9be 100644
--- a/cpp/src/algorithms/TAO/RegressionTaoAdapter.cpp
+++ b/cpp/src/algorithms/TAO/RegressionTaoAdapter.cpp
@@ -156,4 +156,39 @@ void RegressionTaoAdapter::recomputeLeafStats(
}
}
+void RegressionTaoAdapter::refreshNodeBinMetadata(
+ ShapeFunctionNode &node, const std::vector &samples) {
+ if (!node.innerDiscretizer)
+ throw std::runtime_error(
+ "RegressionTaoAdapter: accepted router has no discretizer");
+ arma::Row bins;
+ node.innerDiscretizer->transform(X_, bins);
+ const size_t nBins = node.binToPartition.size();
+ node.binSampleCounts.assign(nBins, 0);
+ node.splitBinWeights.assign(nBins, 0.0);
+ node.splitClassCounts.clear();
+ if (squared_)
+ node.splitLeafStats.assign(
+ nBins, std::vector>(
+ y_.n_rows, std::vector(2, 0.0)));
+ else
+ node.splitLeafStats.clear();
+ for (arma::uword col : samples) {
+ const size_t bin = bins(col);
+ if (bin >= nBins)
+ throw std::runtime_error(
+ "RegressionTaoAdapter: accepted router bin out of range");
+ ++node.binSampleCounts[bin];
+ const double weight = static_cast(w_(col));
+ node.splitBinWeights[bin] += weight;
+ if (!squared_)
+ continue;
+ for (arma::uword o = 0; o < y_.n_rows; ++o) {
+ const double value = static_cast(y_(o, col));
+ node.splitLeafStats[bin][o][0] += weight * value;
+ node.splitLeafStats[bin][o][1] += weight * value * value;
+ }
+ }
+}
+
} // namespace tao
diff --git a/cpp/src/algorithms/TAO/RegressionTaoAdapter.h b/cpp/src/algorithms/TAO/RegressionTaoAdapter.h
index 666c8c5..8e3b1cd 100644
--- a/cpp/src/algorithms/TAO/RegressionTaoAdapter.h
+++ b/cpp/src/algorithms/TAO/RegressionTaoAdapter.h
@@ -52,6 +52,10 @@ class RegressionTaoAdapter final : public ShapeGeneralizedTaoAdapter {
void recomputeLeafStats(
const std::vector> &nodeSamples) override;
+ void refreshNodeBinMetadata(
+ ShapeFunctionNode &node,
+ const std::vector &samples) override;
+
private:
/**
* Per-child negative loss rewards for one sample under the current leaf
diff --git a/cpp/src/algorithms/TAO/TaoAdapter.h b/cpp/src/algorithms/TAO/TaoAdapter.h
index 92f37ec..2c2bccd 100644
--- a/cpp/src/algorithms/TAO/TaoAdapter.h
+++ b/cpp/src/algorithms/TAO/TaoAdapter.h
@@ -123,6 +123,12 @@ class TaoAdapter {
*/
virtual void recomputeLeafStats(
const std::vector> &nodeSamples) = 0;
+
+ /** Refresh exported per-bin metadata for a router accepted by TAO. */
+ virtual void refreshNodeBinMetadata(
+ ShapeFunctionNode &node,
+ const std::vector &samples) = 0;
+
};
} // namespace tao
diff --git a/cpp/src/algorithms/TAO/TaoObjective.cpp b/cpp/src/algorithms/TAO/TaoObjective.cpp
index 9b3d2fa..550b48c 100644
--- a/cpp/src/algorithms/TAO/TaoObjective.cpp
+++ b/cpp/src/algorithms/TAO/TaoObjective.cpp
@@ -5,9 +5,6 @@
#include
#include "algorithms/TAO/TaoObjective.h"
-#include "Discretizers/univariate/UnivariateDiscretizer.h"
-#include "algorithms/missing_values.h"
-
#include
#include
#include
@@ -17,9 +14,9 @@
namespace tao {
TaoObjective::TaoObjective(const NodeCareSet &care, const arma::fmat &X,
- double lambda, double totalSampleWeight)
+ double lambda, double nodeSampleCount)
: care_(care), X_(X), lambda_(lambda),
- totalSampleWeight_(totalSampleWeight), nCare_(care.size()),
+ nodeSampleCount_(nodeSampleCount), nCare_(care.size()),
totalCareWeight_(0.0) {
if (care_.careWeights.empty()) {
totalCareWeight_ = static_cast