sgtlearn 0.2.0__tar.gz → 0.3.0__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (218) hide show
  1. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/PKG-INFO +4 -5
  2. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/README.md +3 -4
  3. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/ROADMAP.md +13 -3
  4. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/CMakeLists.txt +15 -1
  5. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/bindings/Discretizers.cpp +147 -57
  6. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/bindings/ShapeGeneralizedTrees.cpp +38 -13
  7. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/bindings/TreeAlternatingOptimization.cpp +23 -19
  8. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/bindings/_arma_bridge.h +87 -0
  9. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/bindings/_sgt_estimators.h +240 -46
  10. sgtlearn-0.3.0/cpp/src/BranchAssignmentObjectives/AbsoluteErrorBranchAssignment.cpp +214 -0
  11. sgtlearn-0.3.0/cpp/src/BranchAssignmentObjectives/AbsoluteErrorBranchAssignment.h +61 -0
  12. sgtlearn-0.3.0/cpp/src/BranchAssignmentObjectives/AbsoluteErrorBranchAssignmentBst.cpp +125 -0
  13. sgtlearn-0.3.0/cpp/src/BranchAssignmentObjectives/AbsoluteErrorBranchAssignmentBst.h +50 -0
  14. sgtlearn-0.3.0/cpp/src/BranchAssignmentObjectives/AbsoluteErrorBranchAssignmentCommon.h +51 -0
  15. sgtlearn-0.3.0/cpp/src/BranchAssignmentObjectives/AbsoluteErrorBranchAssignmentSort.cpp +128 -0
  16. sgtlearn-0.3.0/cpp/src/BranchAssignmentObjectives/AbsoluteErrorBranchAssignmentSort.h +49 -0
  17. sgtlearn-0.3.0/cpp/src/BranchAssignmentObjectives/BranchAssignmentFactory.cpp +98 -0
  18. sgtlearn-0.3.0/cpp/src/BranchAssignmentObjectives/BranchAssignmentFactory.h +43 -0
  19. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/BranchAssignmentObjectives/BranchAssignmentVariants.h +32 -26
  20. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/BranchAssignmentObjectives/LeafAggregateProcessor.h +22 -14
  21. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/BranchAssignmentObjectives/LeafAggregationBranchAssignment.cpp +4 -13
  22. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/BranchAssignmentObjectives/LeafAggregationBranchAssignment.h +33 -2
  23. sgtlearn-0.3.0/cpp/src/BranchAssignmentObjectives/MaeBranchConfig.h +39 -0
  24. sgtlearn-0.3.0/cpp/src/Criterion.cpp +159 -0
  25. sgtlearn-0.3.0/cpp/src/Criterion.h +52 -0
  26. sgtlearn-0.3.0/cpp/src/Discretizers/ClassificationDiscretizer.h +46 -0
  27. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/Discretizers/InnerDiscretizerBase.h +38 -13
  28. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/Discretizers/RegressionDiscretizer.h +9 -3
  29. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/Discretizers/categorical/CategoricalClassificationDiscretizer.cpp +15 -9
  30. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/Discretizers/categorical/CategoricalClassificationDiscretizer.h +9 -4
  31. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/Discretizers/categorical/CategoricalDiscretizer.h +1 -1
  32. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/Discretizers/categorical/CategoricalDiscretizer.tpp +24 -3
  33. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/Discretizers/categorical/CategoricalRegressionDiscretizer.cpp +3 -3
  34. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/Discretizers/categorical/CategoricalRegressionDiscretizer.h +6 -4
  35. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/Discretizers/factories/DiscretizerFactories.cpp +6 -5
  36. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/Discretizers/factories/DiscretizerFactories.h +6 -4
  37. sgtlearn-0.3.0/cpp/src/Discretizers/pair/PairClassificationDiscretizer.cpp +351 -0
  38. sgtlearn-0.3.0/cpp/src/Discretizers/pair/PairClassificationDiscretizer.h +81 -0
  39. sgtlearn-0.3.0/cpp/src/Discretizers/pair/PairRegressionDiscretizer.cpp +344 -0
  40. sgtlearn-0.3.0/cpp/src/Discretizers/pair/PairRegressionDiscretizer.h +68 -0
  41. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/Discretizers/univariate/UnivariateClassificationDiscretizer.cpp +23 -14
  42. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/Discretizers/univariate/UnivariateClassificationDiscretizer.h +12 -7
  43. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/Discretizers/univariate/UnivariateDiscretizer.h +2 -2
  44. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/Discretizers/univariate/UnivariateRegressionDiscretizer.cpp +17 -14
  45. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/Discretizers/univariate/UnivariateRegressionDiscretizer.h +6 -4
  46. sgtlearn-0.3.0/cpp/src/Estimators/ClassificationShapeGeneralizedTree.cpp +590 -0
  47. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/Estimators/ClassificationShapeGeneralizedTree.h +69 -17
  48. sgtlearn-0.3.0/cpp/src/Estimators/RegressionShapeGeneralizedTree.cpp +688 -0
  49. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/Estimators/RegressionShapeGeneralizedTree.h +40 -16
  50. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/Estimators/ShapeFunctions/NanPartitionRouting.h +83 -99
  51. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/Estimators/ShapeFunctions/ShapeFunctionNode.h +23 -4
  52. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/Estimators/ShapeFunctions/ShapeFunctionSplitSearch.cpp +117 -37
  53. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/Estimators/ShapeFunctions/ShapeFunctionSplitSearch.h +20 -19
  54. sgtlearn-0.3.0/cpp/src/Splitters/categorical/CategoricalClassificationSplitter.cpp +57 -0
  55. sgtlearn-0.3.0/cpp/src/Splitters/categorical/CategoricalClassificationSplitter.h +58 -0
  56. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/Splitters/categorical/CategoricalRegressionSplitter.cpp +26 -21
  57. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/Splitters/categorical/CategoricalRegressionSplitter.h +16 -7
  58. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/Splitters/factories/SplitterFactory.cpp +7 -2
  59. sgtlearn-0.3.0/cpp/src/Splitters/univariate/AbsoluteErrorSplitter.cpp +77 -0
  60. sgtlearn-0.3.0/cpp/src/Splitters/univariate/AbsoluteErrorSplitter.h +42 -0
  61. sgtlearn-0.3.0/cpp/src/Splitters/univariate/ClassificationSplitter.h +90 -0
  62. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/Splitters/univariate/EntropySplitter.h +6 -7
  63. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/Splitters/univariate/GiniSplitter.h +6 -7
  64. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/Splitters/univariate/Splitter.h +1 -1
  65. sgtlearn-0.3.0/cpp/src/Splitters/univariate/SquaredErrorSplitter.cpp +71 -0
  66. sgtlearn-0.3.0/cpp/src/Splitters/univariate/SquaredErrorSplitter.h +42 -0
  67. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/algorithms/CoordinateDescent.h +12 -8
  68. sgtlearn-0.3.0/cpp/src/algorithms/TAO/ClassificationTaoAdapter.cpp +210 -0
  69. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/algorithms/TAO/ClassificationTaoAdapter.h +15 -4
  70. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/algorithms/TAO/RegressionTaoAdapter.cpp +79 -23
  71. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/algorithms/TAO/RegressionTaoAdapter.h +6 -2
  72. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/algorithms/TAO/TaoAdapter.h +6 -0
  73. sgtlearn-0.3.0/cpp/src/algorithms/TAO/TaoObjective.cpp +123 -0
  74. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/algorithms/TAO/TaoObjective.h +15 -24
  75. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/algorithms/TAO/TreeAlternatingOptimization.cpp +71 -27
  76. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/algorithms/TAO/TreeAlternatingOptimization.h +17 -9
  77. sgtlearn-0.3.0/cpp/src/algorithms/WeightedMAETree.cpp +284 -0
  78. sgtlearn-0.3.0/cpp/src/algorithms/WeightedMAETree.h +110 -0
  79. sgtlearn-0.3.0/cpp/tests/bench_mae_branch_assignment.cpp +239 -0
  80. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/tests/test_branch_assignment.cpp +40 -30
  81. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/tests/test_splitters.cpp +23 -23
  82. sgtlearn-0.3.0/cpp/tests/test_weighted_mae_tree.cpp +93 -0
  83. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/docs/api/ensemble.rst +14 -1
  84. sgtlearn-0.3.0/docs/api/estimators.rst +112 -0
  85. sgtlearn-0.3.0/docs/api/plotting.rst +31 -0
  86. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/docs/api/tao.rst +20 -3
  87. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/docs/conf.py +1 -1
  88. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/docs/index.rst +7 -5
  89. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/docs/quickstart.rst +26 -2
  90. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/docs/roadmap.rst +15 -7
  91. sgtlearn-0.3.0/docs/tutorials/bivariate-branching.ipynb +293 -0
  92. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/docs/tutorials/categorical-features.ipynb +2 -2
  93. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/docs/tutorials/feature-importance.ipynb +3 -3
  94. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/docs/tutorials/forests.ipynb +3 -2
  95. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/docs/tutorials/inspecting-trees.ipynb +5 -3
  96. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/docs/tutorials/regression.ipynb +1 -1
  97. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/docs/tutorials/sgt-k.ipynb +1 -1
  98. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/docs/tutorials/shape-functions.ipynb +2 -2
  99. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/docs/tutorials/structure-and-accuracy.ipynb +1 -0
  100. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/docs/tutorials/tao.ipynb +2 -0
  101. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/pyproject.toml +1 -2
  102. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/sgtlearn/_export.py +510 -31
  103. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/sgtlearn/_features.py +3 -2
  104. sgtlearn-0.3.0/sgtlearn/_multioutput.py +142 -0
  105. sgtlearn-0.3.0/sgtlearn/_weights.py +111 -0
  106. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/sgtlearn/base.py +272 -91
  107. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/sgtlearn/datasets.py +1 -3
  108. sgtlearn-0.3.0/sgtlearn/ensemble/__init__.py +4 -0
  109. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/sgtlearn/ensemble/_random_sgforest.py +74 -44
  110. sgtlearn-0.2.0/sgtlearn/ensemble/RandomSGForestClassifier.py → sgtlearn-0.3.0/sgtlearn/ensemble/random_sgforest_classifier.py +90 -31
  111. sgtlearn-0.2.0/sgtlearn/ensemble/RandomSGForestRegressor.py → sgtlearn-0.3.0/sgtlearn/ensemble/random_sgforest_regressor.py +33 -11
  112. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/sgtlearn/tao.py +62 -22
  113. sgtlearn-0.3.0/tests/discretizer_grid.py +11 -0
  114. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/tests/test_categorical_classification_discretizer.py +35 -46
  115. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/tests/test_categorical_regression_discretizer.py +72 -54
  116. sgtlearn-0.3.0/tests/test_datasets.py +44 -0
  117. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/tests/test_features.py +47 -24
  118. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/tests/test_mae_regression_stress.py +14 -43
  119. sgtlearn-0.3.0/tests/test_missing_values.py +209 -0
  120. sgtlearn-0.3.0/tests/test_multioutput.py +77 -0
  121. sgtlearn-0.3.0/tests/test_pairwise_categorical_missing.py +180 -0
  122. sgtlearn-0.3.0/tests/test_pairwise_classifier.py +197 -0
  123. sgtlearn-0.3.0/tests/test_pairwise_regression.py +92 -0
  124. sgtlearn-0.3.0/tests/test_plot_helpers.py +155 -0
  125. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/tests/test_plot_tree.py +128 -65
  126. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/tests/test_random_sgforest_classifier_fidelity.py +38 -12
  127. sgtlearn-0.3.0/tests/test_random_sgforest_pairwise.py +78 -0
  128. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/tests/test_random_sgforest_regressor_fidelity.py +14 -4
  129. sgtlearn-0.3.0/tests/test_random_sgforest_validation.py +134 -0
  130. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/tests/test_sgt_classifier_fidelity.py +33 -8
  131. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/tests/test_sgt_regressor_fidelity.py +16 -6
  132. sgtlearn-0.3.0/tests/test_tao.py +412 -0
  133. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/tests/test_tree_export.py +66 -5
  134. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/tests/test_univariate_classification_discretizer.py +45 -36
  135. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/tests/test_univariate_regression_discretizer.py +77 -76
  136. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/tests/test_weighted_sample.py +119 -16
  137. sgtlearn-0.2.0/cpp/src/BranchAssignmentObjectives/AbsoluteErrorBranchAssignment.cpp +0 -136
  138. sgtlearn-0.2.0/cpp/src/BranchAssignmentObjectives/AbsoluteErrorBranchAssignment.h +0 -53
  139. sgtlearn-0.2.0/cpp/src/BranchAssignmentObjectives/BranchAssignmentFactory.cpp +0 -52
  140. sgtlearn-0.2.0/cpp/src/BranchAssignmentObjectives/BranchAssignmentFactory.h +0 -32
  141. sgtlearn-0.2.0/cpp/src/Criterion.cpp +0 -124
  142. sgtlearn-0.2.0/cpp/src/Criterion.h +0 -36
  143. sgtlearn-0.2.0/cpp/src/Discretizers/ClassificationDiscretizer.h +0 -23
  144. sgtlearn-0.2.0/cpp/src/Estimators/ClassificationShapeGeneralizedTree.cpp +0 -397
  145. sgtlearn-0.2.0/cpp/src/Estimators/RegressionShapeGeneralizedTree.cpp +0 -472
  146. sgtlearn-0.2.0/cpp/src/Splitters/categorical/CategoricalClassificationSplitter.cpp +0 -54
  147. sgtlearn-0.2.0/cpp/src/Splitters/categorical/CategoricalClassificationSplitter.h +0 -45
  148. sgtlearn-0.2.0/cpp/src/Splitters/univariate/AbsoluteErrorSplitter.cpp +0 -59
  149. sgtlearn-0.2.0/cpp/src/Splitters/univariate/AbsoluteErrorSplitter.h +0 -31
  150. sgtlearn-0.2.0/cpp/src/Splitters/univariate/ClassificationSplitter.h +0 -58
  151. sgtlearn-0.2.0/cpp/src/Splitters/univariate/SquaredErrorSplitter.cpp +0 -55
  152. sgtlearn-0.2.0/cpp/src/Splitters/univariate/SquaredErrorSplitter.h +0 -30
  153. sgtlearn-0.2.0/cpp/src/algorithms/TAO/ClassificationTaoAdapter.cpp +0 -125
  154. sgtlearn-0.2.0/cpp/src/algorithms/TAO/TaoObjective.cpp +0 -127
  155. sgtlearn-0.2.0/docs/api/estimators.rst +0 -55
  156. sgtlearn-0.2.0/docs/api/plotting.rst +0 -21
  157. sgtlearn-0.2.0/sgtlearn/_weights.py +0 -64
  158. sgtlearn-0.2.0/sgtlearn/ensemble/__init__.py +0 -4
  159. sgtlearn-0.2.0/tests/discretizer_grid.py +0 -15
  160. sgtlearn-0.2.0/tests/test_missing_values.py +0 -446
  161. sgtlearn-0.2.0/tests/test_plot_helpers.py +0 -545
  162. sgtlearn-0.2.0/tests/test_tao.py +0 -474
  163. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/.github/workflows/workflow.yml +0 -0
  164. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/.gitignore +0 -0
  165. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/.readthedocs.yaml +0 -0
  166. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/LICENSE +0 -0
  167. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/assets/SGT_Viz.png +0 -0
  168. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/README.md +0 -0
  169. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/BranchAssignmentObjectives/BranchAssignment.h +0 -0
  170. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/Discretizers/DiscretizerInputKind.h +0 -0
  171. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/Discretizers/univariate/GainHessianUnivariateDiscretizer.cpp +0 -0
  172. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/Discretizers/univariate/GainHessianUnivariateDiscretizer.h +0 -0
  173. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/Discretizers/univariate/UnivariateDiscretizer.tpp +0 -0
  174. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/Domain/CategoricalSplitCandidate.h +0 -0
  175. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/Domain/FeatureInfo.h +0 -0
  176. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/Domain/LearningCriterion.h +0 -0
  177. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/Domain/LearningFactories.h +0 -0
  178. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/Domain/UnivariateSplitCandidate.h +0 -0
  179. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/Estimators/ShapeFunctions/ShapeFunctionNode.cpp +0 -0
  180. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/Estimators/ShapeGeneralizedTree.cpp +0 -0
  181. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/Estimators/ShapeGeneralizedTree.h +0 -0
  182. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/Splitters/categorical/CategoricalSplitter.h +0 -0
  183. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/Splitters/categorical/CategoricalSplitter.tpp +0 -0
  184. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/Splitters/factories/SplitterFactory.h +0 -0
  185. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/Splitters/univariate/GainHessianSplitter.cpp +0 -0
  186. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/Splitters/univariate/GainHessianSplitter.h +0 -0
  187. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/Splitters/univariate/Splitter.tpp +0 -0
  188. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/algorithms/BinPartitionAssignments.h +0 -0
  189. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/algorithms/FeatureBagging.h +0 -0
  190. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/algorithms/KMeansUtils.h +0 -0
  191. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/algorithms/ShapeBranchingTypes.h +0 -0
  192. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/algorithms/ShapeGeneralizedTreeParams.h +0 -0
  193. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/algorithms/TAO/ShapeGeneralizedTaoAdapter.cpp +0 -0
  194. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/algorithms/TAO/ShapeGeneralizedTaoAdapter.h +0 -0
  195. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/algorithms/TreeBuilder.h +0 -0
  196. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/algorithms/TreeBuilder.tpp +0 -0
  197. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/algorithms/WaveletTreeMAE.cpp +0 -0
  198. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/algorithms/WaveletTreeMAE.h +0 -0
  199. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/algorithms/frontiers.h +0 -0
  200. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/src/algorithms/missing_values.h +0 -0
  201. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/cpp/tests/test_wavelet_tree_mae.cpp +0 -0
  202. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/docs/Makefile +0 -0
  203. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/docs/_static/.gitkeep +0 -0
  204. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/docs/_templates/.gitkeep +0 -0
  205. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/docs/api/datasets.rst +0 -0
  206. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/docs/api/index.rst +0 -0
  207. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/docs/installation.rst +0 -0
  208. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/docs/make.bat +0 -0
  209. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/docs/requirements.txt +0 -0
  210. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/example.py +0 -0
  211. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/examples/plot_tree_demo.py +0 -0
  212. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/setup.py +0 -0
  213. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/sgtlearn/__init__.py +8 -8
  214. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/sgtlearn/py.typed +0 -0
  215. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/tests/__init__.py +0 -0
  216. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/tests/constants.py +0 -0
  217. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/tests/test_max_features_float.py +0 -0
  218. {sgtlearn-0.2.0 → sgtlearn-0.3.0}/tree.png +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.2
2
2
  Name: sgtlearn
3
- Version: 0.2.0
3
+ Version: 0.3.0
4
4
  Summary: Shape Generalized Trees learning library
5
5
  Author-Email: Nakul Upadhya <nakulupadhya1@gmail.com>, Joshua Lee <joshua.lee.9880@gmail.com>, Eldan Cohen <eldan.cohen@utoronto.ca>
6
6
  License: MIT
@@ -31,7 +31,7 @@ Description-Content-Type: text/markdown
31
31
 
32
32
  `sgtlearn` is a Python package for learning [Shape Generalized Trees (SGTs)](https://neurips.cc/virtual/2025/loc/san-diego/poster/115950).
33
33
 
34
- - 🌳 **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.
34
+ - 🌳 **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.
35
35
  - 👁 **Interpretability:** Each node's shape function can be visualized directly.
36
36
  - ⚡ **ShapeCART Algorithm:** An efficient induction method for learning SGTs from data.
37
37
  - 🔀 **Extensions:**
@@ -39,10 +39,9 @@ Description-Content-Type: text/markdown
39
39
  - **SGT<sub>K</sub>:** Multi-way branching generalization.
40
40
  - **Shape²CART & ShapeCART<sub>K</sub>:** Algorithms for learning S²GTs and SGT<sub>K</sub>s.
41
41
 
42
+
42
43
  > [!NOTE]
43
- > 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:
44
- > * Bivariate shape functions (Shape$^2$CART) + Higher branching factors for bivariate splits (Shape$^2$SGT$_K$)
45
- > * Visualization for bivariate splits (ex. contour plots)
44
+ > 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.
46
45
 
47
46
  ## Installation
48
47
 
@@ -4,7 +4,7 @@
4
4
 
5
5
  `sgtlearn` is a Python package for learning [Shape Generalized Trees (SGTs)](https://neurips.cc/virtual/2025/loc/san-diego/poster/115950).
6
6
 
7
- - 🌳 **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.
7
+ - 🌳 **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.
8
8
  - 👁 **Interpretability:** Each node's shape function can be visualized directly.
9
9
  - ⚡ **ShapeCART Algorithm:** An efficient induction method for learning SGTs from data.
10
10
  - 🔀 **Extensions:**
@@ -12,10 +12,9 @@
12
12
  - **SGT<sub>K</sub>:** Multi-way branching generalization.
13
13
  - **Shape²CART & ShapeCART<sub>K</sub>:** Algorithms for learning S²GTs and SGT<sub>K</sub>s.
14
14
 
15
+
15
16
  > [!NOTE]
16
- > 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:
17
- > * Bivariate shape functions (Shape$^2$CART) + Higher branching factors for bivariate splits (Shape$^2$SGT$_K$)
18
- > * Visualization for bivariate splits (ex. contour plots)
17
+ > 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.
19
18
 
20
19
  ## Installation
21
20
 
@@ -17,6 +17,16 @@
17
17
  - [x] **NaN routing at predict**: if training saw missing at that split, follow the stored direction; otherwise route to the majority child.
18
18
 
19
19
  ## v0.3.0
20
- - [ ] multioutput support
21
- - [ ] Shape$^2$CART
22
- - [ ] Shape$^2$CART Random Forest Ensembling
20
+ - [x] multioutput support
21
+ - [x] Opt-in Shape$^2$CART for SGT estimators, including continuous/categorical
22
+ pairs, joint missing routing, and multiway branching ([tutorial](https://sgtlearn.readthedocs.io/en/latest/tutorials/bivariate-branching.html))
23
+ - [x] Shape$^2$CART Random Forest Ensembling
24
+ - [x] Pair-aware TAO refinement ([#48](https://github.com/optimal-uoft/sgtlearn/issues/48))
25
+ - [x] Shape$^2$CART routing heatmap visualization ([#28](https://github.com/optimal-uoft/sgtlearn/issues/28))
26
+
27
+ See the implementation specification in [#42](https://github.com/optimal-uoft/sgtlearn/issues/42)
28
+ and the umbrella issue [#27](https://github.com/optimal-uoft/sgtlearn/issues/27).
29
+
30
+
31
+ ## v1.0.0
32
+ - [ ] Boosting
@@ -77,6 +77,10 @@ add_library(sgtlearn_core STATIC
77
77
  src/Splitters/categorical/CategoricalRegressionSplitter.h
78
78
  src/Splitters/categorical/CategoricalRegressionSplitter.cpp
79
79
  src/Discretizers/univariate/UnivariateClassificationDiscretizer.cpp
80
+ src/Discretizers/pair/PairClassificationDiscretizer.h
81
+ src/Discretizers/pair/PairClassificationDiscretizer.cpp
82
+ src/Discretizers/pair/PairRegressionDiscretizer.h
83
+ src/Discretizers/pair/PairRegressionDiscretizer.cpp
80
84
  src/Splitters/univariate/SquaredErrorSplitter.h
81
85
  src/Splitters/univariate/SquaredErrorSplitter.cpp
82
86
  src/Discretizers/univariate/UnivariateRegressionDiscretizer.cpp
@@ -84,6 +88,8 @@ add_library(sgtlearn_core STATIC
84
88
  src/Discretizers/univariate/GainHessianUnivariateDiscretizer.h
85
89
  src/algorithms/WaveletTreeMAE.h
86
90
  src/algorithms/WaveletTreeMAE.cpp
91
+ src/algorithms/WeightedMAETree.h
92
+ src/algorithms/WeightedMAETree.cpp
87
93
  src/Discretizers/univariate/GainHessianUnivariateDiscretizer.cpp
88
94
  src/Splitters/univariate/AbsoluteErrorSplitter.h
89
95
  src/Splitters/univariate/AbsoluteErrorSplitter.cpp
@@ -119,8 +125,14 @@ add_library(sgtlearn_core STATIC
119
125
  src/Estimators/RegressionShapeGeneralizedTree.h
120
126
  src/Estimators/RegressionShapeGeneralizedTree.cpp
121
127
  src/BranchAssignmentObjectives/BranchAssignment.h
128
+ src/BranchAssignmentObjectives/AbsoluteErrorBranchAssignmentCommon.h
122
129
  src/BranchAssignmentObjectives/AbsoluteErrorBranchAssignment.h
123
130
  src/BranchAssignmentObjectives/AbsoluteErrorBranchAssignment.cpp
131
+ src/BranchAssignmentObjectives/AbsoluteErrorBranchAssignmentBst.h
132
+ src/BranchAssignmentObjectives/AbsoluteErrorBranchAssignmentBst.cpp
133
+ src/BranchAssignmentObjectives/AbsoluteErrorBranchAssignmentSort.h
134
+ src/BranchAssignmentObjectives/AbsoluteErrorBranchAssignmentSort.cpp
135
+ src/BranchAssignmentObjectives/MaeBranchConfig.h
124
136
  src/BranchAssignmentObjectives/LeafAggregateProcessor.h
125
137
  src/BranchAssignmentObjectives/LeafAggregationBranchAssignment.h
126
138
  src/BranchAssignmentObjectives/LeafAggregationBranchAssignment.cpp
@@ -223,8 +235,10 @@ if (SGTLEARN_BUILD_TESTS)
223
235
 
224
236
  add_executable(cpp_tests
225
237
  tests/test_wavelet_tree_mae.cpp
238
+ tests/test_weighted_mae_tree.cpp
226
239
  tests/test_splitters.cpp
227
240
  tests/test_branch_assignment.cpp
241
+ tests/bench_mae_branch_assignment.cpp
228
242
  )
229
243
 
230
244
  target_link_libraries(cpp_tests PRIVATE
@@ -240,4 +254,4 @@ if (SGTLEARN_BUILD_TESTS)
240
254
 
241
255
  endif ()
242
256
 
243
- # endregion
257
+ # endregion
@@ -15,6 +15,9 @@
15
15
  #include <string>
16
16
  #include <variant>
17
17
 
18
+ #define NPY_NO_DEPRECATED_API NPY_1_7_API_VERSION
19
+ #include <numpy/arrayobject.h>
20
+
18
21
  #include "Splitters/univariate/AbsoluteErrorSplitter.h"
19
22
  #include "Discretizers/categorical/CategoricalClassificationDiscretizer.h"
20
23
  #include "Discretizers/categorical/CategoricalRegressionDiscretizer.h"
@@ -58,14 +61,6 @@ py::array_t<size_t> col_to_numpy_1d(const arma::Col<size_t> &col) {
58
61
  return out;
59
62
  }
60
63
 
61
- py::array_t<float> vector_float_to_numpy_1d(const std::vector<float> &v) {
62
- py::array_t<float> out({static_cast<py::ssize_t>(v.size())});
63
- auto buf = out.mutable_unchecked<1>();
64
- for (size_t i = 0; i < v.size(); ++i)
65
- buf(static_cast<py::ssize_t>(i)) = v[i];
66
- return out;
67
- }
68
-
69
64
  std::string normalize_regression_criterion(std::string s) {
70
65
  s = normalize_criterion(std::move(s));
71
66
  if (s == "mse")
@@ -115,6 +110,119 @@ VectorOfVectorsToNumpyList(const std::vector<std::vector<size_t>> &bins) {
115
110
  return out;
116
111
  }
117
112
 
113
+ /**
114
+ * Build a ``(n_outputs, n_samples)`` label matrix from a NumPy array that is
115
+ * either 1-D ``(n_samples,)`` or 2-D ``(n_samples, n_outputs)``.
116
+ */
117
+ arma::Mat<size_t> classification_y_to_arma(const py::array_t<size_t> &y,
118
+ arma::uword n_samples) {
119
+ py::array_t<size_t> ycopy = py::array_t<size_t>::ensure(y);
120
+ arma::Mat<size_t> yArma;
121
+ if (ycopy.ndim() == 1) {
122
+ yArma = carma::arr_to_row<size_t>(ycopy, true);
123
+ } else if (ycopy.ndim() == 2) {
124
+ yArma = arma::Mat<size_t>(carma::arr_to_mat<size_t>(ycopy, true).t());
125
+ } else {
126
+ throw std::invalid_argument(
127
+ "y must be a 1D (n_samples,) or 2D (n_samples, n_outputs) numpy array");
128
+ }
129
+ if (yArma.n_cols != n_samples)
130
+ throw std::invalid_argument("y length must match X.shape[0]");
131
+ return yArma;
132
+ }
133
+
134
+ /**
135
+ * Build a ``(n_outputs, n_samples)`` target matrix from a NumPy array that is
136
+ * either 1-D ``(n_samples,)`` or 2-D ``(n_samples, n_outputs)``.
137
+ */
138
+ arma::Mat<float> regression_y_to_arma(const py::array_t<float> &y,
139
+ arma::uword n_samples) {
140
+ py::array_t<float> ycopy = py::array_t<float>::ensure(y);
141
+ arma::Mat<float> yArma;
142
+ if (ycopy.ndim() == 1) {
143
+ yArma = carma::arr_to_row<float>(ycopy, true);
144
+ } else if (ycopy.ndim() == 2) {
145
+ yArma = arma::Mat<float>(carma::arr_to_mat<float>(ycopy, true).t());
146
+ } else {
147
+ throw std::invalid_argument(
148
+ "y must be a 1D (n_samples,) or 2D (n_samples, n_outputs) numpy array");
149
+ }
150
+ if (yArma.n_cols != n_samples)
151
+ throw std::invalid_argument("y length must match X.shape[0]");
152
+ return yArma;
153
+ }
154
+
155
+ /**
156
+ * Resolve a Python ``numClasses`` argument (int or sequence of ints) into a
157
+ * per-output class-count vector of length ``n_outputs``. An int is broadcast
158
+ * to every output; a sequence must match ``n_outputs``.
159
+ */
160
+ std::vector<size_t> num_classes_per_output_from_py(const py::object &numClasses,
161
+ arma::uword n_outputs) {
162
+ py::object numbers = py::module_::import("numbers");
163
+ if (py::isinstance<py::bool_>(numClasses))
164
+ throw std::invalid_argument("numClasses cannot be bool");
165
+ if (py::isinstance(numClasses, numbers.attr("Integral")))
166
+ return std::vector<size_t>(static_cast<size_t>(n_outputs),
167
+ py::cast<size_t>(numClasses));
168
+ if (py::isinstance<py::iterable>(numClasses)) {
169
+ std::vector<size_t> out;
170
+ for (const py::handle item : numClasses)
171
+ out.push_back(py::cast<size_t>(py::reinterpret_borrow<py::object>(item)));
172
+ if (out.size() != static_cast<size_t>(n_outputs))
173
+ throw std::invalid_argument(
174
+ "numClasses sequence length must match the number of outputs "
175
+ "(y.shape[1])");
176
+ return out;
177
+ }
178
+ throw std::invalid_argument(
179
+ "numClasses must be an int or a sequence of ints");
180
+ }
181
+
182
+ /**
183
+ * Convert per-bin classification predictions (one vector of length
184
+ * ``n_outputs`` per bin) into a NumPy array: 1-D ``(n_bins,)`` for a single
185
+ * output, else 2-D ``(n_bins, n_outputs)``.
186
+ */
187
+ template <typename T>
188
+ py::array bin_predictions_to_numpy_impl(const std::vector<std::vector<T>> &preds) {
189
+ const size_t nBins = preds.size();
190
+ // Bins with no samples (e.g. the trailing NaN routing bin) carry an empty
191
+ // prediction vector, so derive the output width from the widest bin.
192
+ size_t nOut = 0;
193
+ for (const auto &p : preds)
194
+ nOut = std::max(nOut, p.size());
195
+ if (nOut == 0)
196
+ nOut = 1;
197
+ if (nOut == 1) {
198
+ py::array_t<T> out({static_cast<py::ssize_t>(nBins)});
199
+ auto buf = out.template mutable_unchecked<1>();
200
+ for (size_t i = 0; i < nBins; ++i)
201
+ buf(static_cast<py::ssize_t>(i)) =
202
+ preds[i].empty() ? static_cast<T>(0) : preds[i][0];
203
+ return out;
204
+ }
205
+ py::array_t<T> out(
206
+ {static_cast<py::ssize_t>(nBins), static_cast<py::ssize_t>(nOut)});
207
+ auto buf = out.template mutable_unchecked<2>();
208
+ for (size_t i = 0; i < nBins; ++i)
209
+ for (size_t o = 0; o < nOut; ++o)
210
+ buf(static_cast<py::ssize_t>(i), static_cast<py::ssize_t>(o)) =
211
+ o < preds[i].size() ? preds[i][o] : static_cast<T>(0);
212
+ return out;
213
+ }
214
+
215
+ py::array
216
+ bin_predictions_to_numpy(const std::vector<std::vector<size_t>> &preds) {
217
+ return bin_predictions_to_numpy_impl(preds);
218
+ }
219
+
220
+ /** Regression counterpart of :func:`bin_predictions_to_numpy`. */
221
+ py::array
222
+ bin_predictions_to_numpy(const std::vector<std::vector<float>> &preds) {
223
+ return bin_predictions_to_numpy_impl(preds);
224
+ }
225
+
118
226
  class UnivariateClassificationDiscretizerPy {
119
227
  std::variant<GiniDisc, EntropyDisc> impl_;
120
228
 
@@ -132,7 +240,7 @@ public:
132
240
  }
133
241
 
134
242
  void Train(const py::array_t<float> &X, const py::array_t<size_t> &features,
135
- const py::array_t<size_t> &y, size_t numClasses,
243
+ const py::array_t<size_t> &y, py::object numClasses,
136
244
  size_t minLeafSize, double minGainSplit, size_t maxDepth,
137
245
  size_t maxLeafNodes, py::object sample_weights = py::none()) {
138
246
  py::array_t<float> Xcopy = py::array_t<float>::ensure(X);
@@ -148,12 +256,9 @@ public:
148
256
  const arma::uvec featuresArma =
149
257
  arma::conv_to<arma::uvec>::from(carma::arr_to_col<size_t>(fcopy, true));
150
258
 
151
- py::array_t<size_t> ycopy = py::array_t<size_t>::ensure(y);
152
- if (ycopy.ndim() != 1)
153
- throw std::invalid_argument("y must be a 1D numpy array");
154
- arma::Mat<size_t> yArma = carma::arr_to_row<size_t>(ycopy, true);
155
- if (yArma.n_elem != xArma.n_cols)
156
- throw std::invalid_argument("y length must match X.shape[0]");
259
+ arma::Mat<size_t> yArma = classification_y_to_arma(y, xArma.n_cols);
260
+ const std::vector<size_t> nClassesPerOutput =
261
+ num_classes_per_output_from_py(numClasses, yArma.n_rows);
157
262
 
158
263
  const arma::Row<float> wRow =
159
264
  sampleWeightRowFromPy(sample_weights, xArma.n_cols);
@@ -161,7 +266,7 @@ public:
161
266
  arma::uvec featuresMut = featuresArma;
162
267
  std::visit(
163
268
  [&](auto &d) {
164
- d.Train(xArma, featuresMut, yArma, numClasses, minLeafSize,
269
+ d.Train(xArma, featuresMut, yArma, nClassesPerOutput, minLeafSize,
165
270
  minGainSplit, maxDepth, maxLeafNodes, wRow);
166
271
  },
167
272
  impl_);
@@ -198,15 +303,9 @@ public:
198
303
  impl_);
199
304
  }
200
305
 
201
- py::array_t<size_t> getBinPredictions() {
306
+ py::array getBinPredictions() {
202
307
  return std::visit(
203
- [](auto &d) {
204
- const auto &preds = d.getBinPredictions();
205
- arma::Col<size_t> col(preds.size());
206
- for (size_t i = 0; i < preds.size(); ++i)
207
- col(i) = preds[i];
208
- return col_to_numpy_1d(col);
209
- },
308
+ [](auto &d) { return bin_predictions_to_numpy(d.getBinPredictions()); },
210
309
  impl_);
211
310
  }
212
311
 
@@ -252,12 +351,7 @@ public:
252
351
  const arma::uvec featuresArma =
253
352
  arma::conv_to<arma::uvec>::from(carma::arr_to_col<size_t>(fcopy, true));
254
353
 
255
- py::array_t<float> ycopy = py::array_t<float>::ensure(y);
256
- if (ycopy.ndim() != 1)
257
- throw std::invalid_argument("y must be a 1D numpy array");
258
- arma::Mat<float> yArma = carma::arr_to_row<float>(ycopy, true);
259
- if (yArma.n_elem != xArma.n_cols)
260
- throw std::invalid_argument("y length must match X.shape[0]");
354
+ arma::Mat<float> yArma = regression_y_to_arma(y, xArma.n_cols);
261
355
 
262
356
  const arma::Row<float> wRow =
263
357
  sampleWeightRowFromPy(sample_weights, xArma.n_cols);
@@ -302,11 +396,9 @@ public:
302
396
  impl_);
303
397
  }
304
398
 
305
- py::array_t<float> getBinPredictions() {
399
+ py::array getBinPredictions() {
306
400
  return std::visit(
307
- [](auto &d) {
308
- return vector_float_to_numpy_1d(d.getBinPredictions());
309
- },
401
+ [](auto &d) { return bin_predictions_to_numpy(d.getBinPredictions()); },
310
402
  impl_);
311
403
  }
312
404
 
@@ -336,7 +428,7 @@ public:
336
428
  }
337
429
 
338
430
  void Train(const py::array_t<float> &X, const py::array_t<size_t> &features,
339
- const py::array_t<size_t> &y, size_t numClasses,
431
+ const py::array_t<size_t> &y, py::object numClasses,
340
432
  size_t minLeafSize, double minGainSplit, size_t maxDepth,
341
433
  size_t maxLeafNodes, py::object sample_weights = py::none()) {
342
434
  py::array_t<float> Xcopy = py::array_t<float>::ensure(X);
@@ -352,19 +444,16 @@ public:
352
444
  const arma::uvec featuresArma =
353
445
  arma::conv_to<arma::uvec>::from(carma::arr_to_col<size_t>(fcopy, true));
354
446
 
355
- py::array_t<size_t> ycopy = py::array_t<size_t>::ensure(y);
356
- if (ycopy.ndim() != 1)
357
- throw std::invalid_argument("y must be a 1D numpy array");
358
- arma::Mat<size_t> yArma = carma::arr_to_row<size_t>(ycopy, true);
359
- if (yArma.n_elem != xArma.n_cols)
360
- throw std::invalid_argument("y length must match X.shape[0]");
447
+ arma::Mat<size_t> yArma = classification_y_to_arma(y, xArma.n_cols);
448
+ const std::vector<size_t> nClassesPerOutput =
449
+ num_classes_per_output_from_py(numClasses, yArma.n_rows);
361
450
 
362
451
  const arma::Row<float> wRow =
363
452
  sampleWeightRowFromPy(sample_weights, xArma.n_cols);
364
453
 
365
454
  arma::uvec featuresMut = featuresArma;
366
- impl_.Train(xArma, featuresMut, yArma, numClasses, minLeafSize, minGainSplit,
367
- maxDepth, maxLeafNodes, wRow);
455
+ impl_.Train(xArma, featuresMut, yArma, nClassesPerOutput, minLeafSize,
456
+ minGainSplit, maxDepth, maxLeafNodes, wRow);
368
457
  }
369
458
 
370
459
  py::array_t<size_t> transform(const py::array_t<float> &X) {
@@ -390,12 +479,8 @@ public:
390
479
  return VectorOfVectorsToNumpyList(impl_.inSampleDiscretizations());
391
480
  }
392
481
 
393
- py::array_t<size_t> getBinPredictions() {
394
- const auto &preds = impl_.getBinPredictions();
395
- arma::Col<size_t> col(preds.size());
396
- for (size_t i = 0; i < preds.size(); ++i)
397
- col(i) = preds[i];
398
- return col_to_numpy_1d(col);
482
+ py::array getBinPredictions() {
483
+ return bin_predictions_to_numpy(impl_.getBinPredictions());
399
484
  }
400
485
 
401
486
  size_t getNumLeaves() const { return impl_.numLeaves(); }
@@ -436,12 +521,7 @@ public:
436
521
  const arma::uvec featuresArma =
437
522
  arma::conv_to<arma::uvec>::from(carma::arr_to_col<size_t>(fcopy, true));
438
523
 
439
- py::array_t<float> ycopy = py::array_t<float>::ensure(y);
440
- if (ycopy.ndim() != 1)
441
- throw std::invalid_argument("y must be a 1D numpy array");
442
- arma::Mat<float> yArma = carma::arr_to_row<float>(ycopy, true);
443
- if (yArma.n_elem != xArma.n_cols)
444
- throw std::invalid_argument("y length must match X.shape[0]");
524
+ arma::Mat<float> yArma = regression_y_to_arma(y, xArma.n_cols);
445
525
 
446
526
  const arma::Row<float> wRow =
447
527
  sampleWeightRowFromPy(sample_weights, xArma.n_cols);
@@ -474,8 +554,8 @@ public:
474
554
  return VectorOfVectorsToNumpyList(impl_.inSampleDiscretizations());
475
555
  }
476
556
 
477
- py::array_t<float> getBinPredictions() {
478
- return vector_float_to_numpy_1d(impl_.getBinPredictions());
557
+ py::array getBinPredictions() {
558
+ return bin_predictions_to_numpy(impl_.getBinPredictions());
479
559
  }
480
560
 
481
561
  size_t getNumLeaves() const { return impl_.numLeaves(); }
@@ -486,6 +566,16 @@ public:
486
566
  } // namespace
487
567
 
488
568
  PYBIND11_MODULE(Discretizers, m) {
569
+ // CARMA's allocator may lazily call _import_array() on free/alloc; prime it
570
+ // here (GIL held) so destructions during ShapeGeneralizedTrees.fit (which can
571
+ // resolve to this module's statically-linked core symbols) stay safe.
572
+ if (_import_array() < 0) {
573
+ PyErr_Clear();
574
+ throw std::runtime_error(
575
+ "Discretizers: numpy.core.multiarray failed to import; "
576
+ "ensure numpy is installed and importable before importing this module");
577
+ }
578
+
489
579
  py::class_<UnivariateClassificationDiscretizerPy>(
490
580
  m, "UnivariateClassificationDiscretizer")
491
581
  .def(py::init<std::string>(), py::arg("criterion") = "gini")
@@ -43,7 +43,7 @@ PYBIND11_MODULE(ShapeGeneralizedTrees, m) {
43
43
 
44
44
  py::class_<ClassificationShapeGeneralizedTreePy>(
45
45
  m, "ClassificationShapeGeneralizedTree")
46
- .def(py::init([](std::string criterion, size_t num_classes,
46
+ .def(py::init([](std::string criterion, py::object num_classes,
47
47
  size_t num_partitions, size_t outer_min_leaf_size,
48
48
  double outer_min_gain_split, size_t outer_max_depth,
49
49
  size_t outer_max_leaf_nodes, size_t inner_min_leaf_size,
@@ -52,15 +52,17 @@ PYBIND11_MODULE(ShapeGeneralizedTrees, m) {
52
52
  size_t coordinate_descent_max_iters,
53
53
  size_t coordinate_descent_patience,
54
54
  bool coordinate_descent_smart_init, uint64_t random_state,
55
- py::object max_features) {
55
+ py::object max_features, size_t pairwise_candidates,
56
+ double pairwise_penalty) {
56
57
  return ClassificationShapeGeneralizedTreePy(
57
- std::move(criterion), num_classes, num_partitions,
58
+ std::move(criterion), std::move(num_classes), num_partitions,
58
59
  outer_min_leaf_size, outer_min_gain_split, outer_max_depth,
59
60
  outer_max_leaf_nodes, inner_min_leaf_size, inner_min_gain_split,
60
61
  inner_max_depth, inner_max_leaf_nodes,
61
62
  coordinate_descent_max_iters, coordinate_descent_patience,
62
63
  coordinate_descent_smart_init, random_state,
63
- std::move(max_features));
64
+ std::move(max_features), pairwise_candidates,
65
+ pairwise_penalty);
64
66
  }),
65
67
  py::arg("criterion") = "gini", py::arg("num_classes"),
66
68
  py::arg("num_partitions") = 2,
@@ -76,25 +78,38 @@ PYBIND11_MODULE(ShapeGeneralizedTrees, m) {
76
78
  py::arg("coordinate_descent_patience") = 5,
77
79
  py::arg("coordinate_descent_smart_init") = true,
78
80
  py::arg("random_state") = 42,
79
- py::arg("max_features") = py::none())
81
+ py::arg("max_features") = py::none(),
82
+ py::arg("pairwise_candidates") = 0,
83
+ py::arg("pairwise_penalty") = 0.0)
80
84
  .def("fit", &ClassificationShapeGeneralizedTreePy::fit, py::arg("X"),
81
85
  py::arg("y"), py::arg("sample_weight") = py::none(),
82
86
  py::arg("features"),
83
87
  "Fit the routing tree. X is (n_samples, n_features) float32; y is "
84
- "1-D uint class labels. Optional sample_weight is 1-D float32.")
88
+ "uint class labels, 1-D (n_samples,) or 2-D (n_samples, n_outputs). "
89
+ "Optional sample_weight is 1-D float32.")
85
90
  .def("predict", &ClassificationShapeGeneralizedTreePy::predict,
86
91
  py::arg("X"),
87
- "Predict class labels for X (shape (n_samples, n_features)).")
92
+ "Predict class labels for X (shape (n_samples, n_features)). Returns "
93
+ "(n_samples,) for a single output or (n_samples, n_outputs) for "
94
+ "multi-output.")
88
95
  .def("predict_proba", &ClassificationShapeGeneralizedTreePy::predictProba,
89
96
  py::arg("X"),
90
- "Predict class probabilities for X; output shape "
91
- "(n_samples, n_classes).")
97
+ "Predict class probabilities for X. Single output: a "
98
+ "(n_samples, n_classes) array. Multi-output: a list of such arrays, "
99
+ "one per output.")
92
100
  .def_property_readonly(
93
101
  "num_leaves", &ClassificationShapeGeneralizedTreePy::numLeaves)
94
102
  .def_property_readonly(
95
103
  "num_nodes", &ClassificationShapeGeneralizedTreePy::numNodes)
104
+ .def_property_readonly(
105
+ "n_outputs", &ClassificationShapeGeneralizedTreePy::nOutputs)
106
+ .def_property_readonly(
107
+ "classes_per_output",
108
+ &ClassificationShapeGeneralizedTreePy::classesPerOutput)
96
109
  .def_property_readonly(
97
110
  "is_fitted", &ClassificationShapeGeneralizedTreePy::isFitted)
111
+ .def_property_readonly(
112
+ "has_pair_nodes", &ClassificationShapeGeneralizedTreePy::hasPairNodes)
98
113
  .def_property_readonly(
99
114
  "feature_importance",
100
115
  &ClassificationShapeGeneralizedTreePy::featureImportance,
@@ -113,14 +128,16 @@ PYBIND11_MODULE(ShapeGeneralizedTrees, m) {
113
128
  size_t coordinate_descent_max_iters,
114
129
  size_t coordinate_descent_patience,
115
130
  bool coordinate_descent_smart_init, uint64_t random_state,
116
- py::object max_features) {
131
+ py::object max_features, size_t pairwise_candidates,
132
+ double pairwise_penalty) {
117
133
  return RegressionShapeGeneralizedTreePy(
118
134
  std::move(criterion), num_partitions, outer_min_leaf_size,
119
135
  outer_min_gain_split, outer_max_depth, outer_max_leaf_nodes,
120
136
  inner_min_leaf_size, inner_min_gain_split, inner_max_depth,
121
137
  inner_max_leaf_nodes, coordinate_descent_max_iters,
122
138
  coordinate_descent_patience, coordinate_descent_smart_init,
123
- random_state, std::move(max_features));
139
+ random_state, std::move(max_features), pairwise_candidates,
140
+ pairwise_penalty);
124
141
  }),
125
142
  py::arg("criterion") = "squared_error",
126
143
  py::arg("num_partitions") = 2, py::arg("outer_min_leaf_size") = 1,
@@ -133,6 +150,8 @@ PYBIND11_MODULE(ShapeGeneralizedTrees, m) {
133
150
  py::arg("coordinate_descent_patience") = 5,
134
151
  py::arg("coordinate_descent_smart_init") = true,
135
152
  py::arg("random_state") = 42, py::arg("max_features") = py::none(),
153
+ py::arg("pairwise_candidates") = 0,
154
+ py::arg("pairwise_penalty") = 0.0,
136
155
  R"(Regression tree: inner bins are round-robin seeded. ``squared_error`` runs
137
156
  coordinate descent and keeps the map only if branch MSE improves clearly vs the seed;
138
157
  otherwise the snapshot is restored and the branch objective is rebuilt.
@@ -142,15 +161,21 @@ is accepted for API parity with ClassificationShapeGeneralizedTree but ignored.)
142
161
  py::arg("y"), py::arg("sample_weight") = py::none(),
143
162
  py::arg("features"),
144
163
  "Fit the routing tree. X is (n_samples, n_features) float32; y is "
145
- "1-D float32 targets. Optional sample_weight is 1-D float32.")
164
+ "float32 targets, 1-D (n_samples,) or 2-D (n_samples, n_outputs). "
165
+ "Optional sample_weight is 1-D float32.")
146
166
  .def("predict", &RegressionShapeGeneralizedTreePy::predict, py::arg("X"),
147
- "Predict scalar targets for X (shape (n_samples,)).")
167
+ "Predict targets for X. Returns (n_samples,) for a single output or "
168
+ "(n_samples, n_outputs) for multi-output.")
148
169
  .def_property_readonly("num_leaves",
149
170
  &RegressionShapeGeneralizedTreePy::numLeaves)
150
171
  .def_property_readonly("num_nodes",
151
172
  &RegressionShapeGeneralizedTreePy::numNodes)
173
+ .def_property_readonly("n_outputs",
174
+ &RegressionShapeGeneralizedTreePy::nOutputs)
152
175
  .def_property_readonly("is_fitted",
153
176
  &RegressionShapeGeneralizedTreePy::isFitted)
177
+ .def_property_readonly(
178
+ "has_pair_nodes", &RegressionShapeGeneralizedTreePy::hasPairNodes)
154
179
  .def_property_readonly(
155
180
  "feature_importance",
156
181
  &RegressionShapeGeneralizedTreePy::featureImportance,