sgtlearn 0.2.0__tar.gz → 0.3.1__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 (280) hide show
  1. sgtlearn-0.3.1/CONTEXT.md +46 -0
  2. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/PKG-INFO +20 -5
  3. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/README.md +19 -4
  4. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/ROADMAP.md +15 -3
  5. sgtlearn-0.3.1/benchmarks/check_outer_growth.py +23 -0
  6. sgtlearn-0.3.1/benchmarks/compare_outer_growth.py +155 -0
  7. sgtlearn-0.3.1/benchmarks/outer_growth.md +102 -0
  8. sgtlearn-0.3.1/benchmarks/outer_growth.py +261 -0
  9. sgtlearn-0.3.1/benchmarks/outer_growth_manifest.json +199 -0
  10. sgtlearn-0.3.1/benchmarks/results/outer-growth-baseline/build-provenance.txt +79 -0
  11. sgtlearn-0.3.1/benchmarks/results/outer-growth-baseline/dependencies.txt +28 -0
  12. sgtlearn-0.3.1/benchmarks/results/outer-growth-baseline/manifest.json +199 -0
  13. sgtlearn-0.3.1/benchmarks/results/outer-growth-baseline/raw.jsonl +99 -0
  14. sgtlearn-0.3.1/benchmarks/results/outer-growth-baseline/runner.sha256 +1 -0
  15. sgtlearn-0.3.1/benchmarks/results/outer-growth-baseline/summary.json +1154 -0
  16. sgtlearn-0.3.1/benchmarks/results/outer-growth-baseline/summary.md +25 -0
  17. sgtlearn-0.3.1/benchmarks/results/outer-growth-branching-entropy/comparison.json +391 -0
  18. sgtlearn-0.3.1/benchmarks/results/outer-growth-branching-entropy/comparison.md +19 -0
  19. sgtlearn-0.3.1/benchmarks/results/outer-growth-branching-entropy/raw.jsonl +22 -0
  20. sgtlearn-0.3.1/benchmarks/results/outer-growth-branching-entropy/settings.json +209 -0
  21. sgtlearn-0.3.1/benchmarks/results/outer-growth-branching-mse/comparison.json +415 -0
  22. sgtlearn-0.3.1/benchmarks/results/outer-growth-branching-mse/comparison.md +19 -0
  23. sgtlearn-0.3.1/benchmarks/results/outer-growth-branching-mse/raw.jsonl +22 -0
  24. sgtlearn-0.3.1/benchmarks/results/outer-growth-branching-mse/settings.json +209 -0
  25. sgtlearn-0.3.1/benchmarks/results/outer-growth-candidate-build/README.md +26 -0
  26. sgtlearn-0.3.1/benchmarks/results/outer-growth-candidate-build/build-provenance.txt +83 -0
  27. sgtlearn-0.3.1/benchmarks/results/outer-growth-candidate-build/dependencies.txt +28 -0
  28. sgtlearn-0.3.1/benchmarks/results/outer-growth-candidate-build/imports.json +49 -0
  29. sgtlearn-0.3.1/benchmarks/results/outer-growth-comparison/comparison.json +4511 -0
  30. sgtlearn-0.3.1/benchmarks/results/outer-growth-comparison/comparison.md +49 -0
  31. sgtlearn-0.3.1/benchmarks/results/outer-growth-comparison/raw.jsonl +198 -0
  32. sgtlearn-0.3.1/benchmarks/results/outer-growth-comparison/settings.json +209 -0
  33. sgtlearn-0.3.1/benchmarks/results/outer-growth-confirmation/comparison.json +2979 -0
  34. sgtlearn-0.3.1/benchmarks/results/outer-growth-confirmation/comparison.md +33 -0
  35. sgtlearn-0.3.1/benchmarks/results/outer-growth-confirmation/raw.jsonl +110 -0
  36. sgtlearn-0.3.1/benchmarks/results/outer-growth-confirmation/settings.json +209 -0
  37. sgtlearn-0.3.1/benchmarks/results/outer-growth-final.md +126 -0
  38. sgtlearn-0.3.1/benchmarks/results/outer-growth-verification.md +105 -0
  39. sgtlearn-0.3.1/benchmarks/results/outer-growth-work-diagnostic/README.md +23 -0
  40. sgtlearn-0.3.1/benchmarks/results/outer-growth-work-diagnostic/baseline.json +1 -0
  41. sgtlearn-0.3.1/benchmarks/results/outer-growth-work-diagnostic/candidate.json +1 -0
  42. sgtlearn-0.3.1/benchmarks/results/outer-growth-work-diagnostic/diagnostic.py +57 -0
  43. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/CMakeLists.txt +17 -1
  44. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/bindings/Discretizers.cpp +147 -57
  45. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/bindings/ShapeGeneralizedTrees.cpp +45 -21
  46. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/bindings/TreeAlternatingOptimization.cpp +23 -19
  47. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/bindings/_arma_bridge.h +87 -0
  48. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/bindings/_sgt_estimators.h +244 -52
  49. sgtlearn-0.3.1/cpp/src/BranchAssignmentObjectives/AbsoluteErrorBranchAssignment.cpp +214 -0
  50. sgtlearn-0.3.1/cpp/src/BranchAssignmentObjectives/AbsoluteErrorBranchAssignment.h +61 -0
  51. sgtlearn-0.3.1/cpp/src/BranchAssignmentObjectives/AbsoluteErrorBranchAssignmentBst.cpp +125 -0
  52. sgtlearn-0.3.1/cpp/src/BranchAssignmentObjectives/AbsoluteErrorBranchAssignmentBst.h +50 -0
  53. sgtlearn-0.3.1/cpp/src/BranchAssignmentObjectives/AbsoluteErrorBranchAssignmentCommon.h +51 -0
  54. sgtlearn-0.3.1/cpp/src/BranchAssignmentObjectives/AbsoluteErrorBranchAssignmentSort.cpp +128 -0
  55. sgtlearn-0.3.1/cpp/src/BranchAssignmentObjectives/AbsoluteErrorBranchAssignmentSort.h +49 -0
  56. sgtlearn-0.3.1/cpp/src/BranchAssignmentObjectives/BranchAssignmentFactory.cpp +98 -0
  57. sgtlearn-0.3.1/cpp/src/BranchAssignmentObjectives/BranchAssignmentFactory.h +43 -0
  58. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/BranchAssignmentObjectives/BranchAssignmentVariants.h +32 -26
  59. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/BranchAssignmentObjectives/LeafAggregateProcessor.h +22 -14
  60. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/BranchAssignmentObjectives/LeafAggregationBranchAssignment.cpp +4 -13
  61. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/BranchAssignmentObjectives/LeafAggregationBranchAssignment.h +33 -2
  62. sgtlearn-0.3.1/cpp/src/BranchAssignmentObjectives/MaeBranchConfig.h +39 -0
  63. sgtlearn-0.3.1/cpp/src/Criterion.cpp +159 -0
  64. sgtlearn-0.3.1/cpp/src/Criterion.h +52 -0
  65. sgtlearn-0.3.1/cpp/src/Discretizers/ClassificationDiscretizer.h +46 -0
  66. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Discretizers/InnerDiscretizerBase.h +43 -13
  67. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Discretizers/RegressionDiscretizer.h +9 -3
  68. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Discretizers/categorical/CategoricalClassificationDiscretizer.cpp +16 -10
  69. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Discretizers/categorical/CategoricalClassificationDiscretizer.h +9 -4
  70. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Discretizers/categorical/CategoricalDiscretizer.h +19 -3
  71. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Discretizers/categorical/CategoricalDiscretizer.tpp +72 -11
  72. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Discretizers/categorical/CategoricalRegressionDiscretizer.cpp +4 -4
  73. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Discretizers/categorical/CategoricalRegressionDiscretizer.h +6 -4
  74. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Discretizers/factories/DiscretizerFactories.cpp +6 -5
  75. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Discretizers/factories/DiscretizerFactories.h +6 -4
  76. sgtlearn-0.3.1/cpp/src/Discretizers/pair/PairClassificationDiscretizer.cpp +356 -0
  77. sgtlearn-0.3.1/cpp/src/Discretizers/pair/PairClassificationDiscretizer.h +106 -0
  78. sgtlearn-0.3.1/cpp/src/Discretizers/pair/PairRegressionDiscretizer.cpp +349 -0
  79. sgtlearn-0.3.1/cpp/src/Discretizers/pair/PairRegressionDiscretizer.h +69 -0
  80. sgtlearn-0.3.1/cpp/src/Discretizers/univariate/NumericFallbackDiscretizer.cpp +233 -0
  81. sgtlearn-0.3.1/cpp/src/Discretizers/univariate/NumericFallbackDiscretizer.h +19 -0
  82. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Discretizers/univariate/UnivariateClassificationDiscretizer.cpp +23 -14
  83. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Discretizers/univariate/UnivariateClassificationDiscretizer.h +12 -7
  84. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Discretizers/univariate/UnivariateDiscretizer.h +13 -2
  85. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Discretizers/univariate/UnivariateDiscretizer.tpp +3 -0
  86. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Discretizers/univariate/UnivariateRegressionDiscretizer.cpp +17 -14
  87. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Discretizers/univariate/UnivariateRegressionDiscretizer.h +6 -4
  88. sgtlearn-0.3.1/cpp/src/Estimators/ClassificationShapeGeneralizedTree.cpp +616 -0
  89. sgtlearn-0.3.1/cpp/src/Estimators/ClassificationShapeGeneralizedTree.h +192 -0
  90. sgtlearn-0.3.1/cpp/src/Estimators/RegressionShapeGeneralizedTree.cpp +715 -0
  91. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Estimators/RegressionShapeGeneralizedTree.h +45 -21
  92. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Estimators/ShapeFunctions/NanPartitionRouting.h +83 -99
  93. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Estimators/ShapeFunctions/ShapeFunctionNode.h +33 -7
  94. sgtlearn-0.3.1/cpp/src/Estimators/ShapeFunctions/ShapeFunctionSplitSearch.cpp +402 -0
  95. sgtlearn-0.3.1/cpp/src/Estimators/ShapeFunctions/ShapeFunctionSplitSearch.h +120 -0
  96. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Estimators/ShapeGeneralizedTree.cpp +7 -1
  97. sgtlearn-0.3.1/cpp/src/Splitters/categorical/CategoricalClassificationSplitter.cpp +57 -0
  98. sgtlearn-0.3.1/cpp/src/Splitters/categorical/CategoricalClassificationSplitter.h +58 -0
  99. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Splitters/categorical/CategoricalRegressionSplitter.cpp +26 -21
  100. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Splitters/categorical/CategoricalRegressionSplitter.h +16 -7
  101. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Splitters/factories/SplitterFactory.cpp +7 -2
  102. sgtlearn-0.3.1/cpp/src/Splitters/univariate/AbsoluteErrorSplitter.cpp +77 -0
  103. sgtlearn-0.3.1/cpp/src/Splitters/univariate/AbsoluteErrorSplitter.h +42 -0
  104. sgtlearn-0.3.1/cpp/src/Splitters/univariate/ClassificationSplitter.h +90 -0
  105. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Splitters/univariate/EntropySplitter.h +6 -7
  106. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Splitters/univariate/GiniSplitter.h +6 -7
  107. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Splitters/univariate/Splitter.h +1 -1
  108. sgtlearn-0.3.1/cpp/src/Splitters/univariate/SquaredErrorSplitter.cpp +71 -0
  109. sgtlearn-0.3.1/cpp/src/Splitters/univariate/SquaredErrorSplitter.h +42 -0
  110. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/algorithms/BinPartitionAssignments.h +3 -10
  111. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/algorithms/CoordinateDescent.h +27 -8
  112. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/algorithms/KMeansUtils.h +13 -11
  113. sgtlearn-0.3.1/cpp/src/algorithms/OuterTreeBuilder.h +75 -0
  114. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/algorithms/ShapeGeneralizedTreeParams.h +4 -10
  115. sgtlearn-0.3.1/cpp/src/algorithms/TAO/ClassificationTaoAdapter.cpp +210 -0
  116. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/algorithms/TAO/ClassificationTaoAdapter.h +15 -4
  117. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/algorithms/TAO/RegressionTaoAdapter.cpp +79 -23
  118. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/algorithms/TAO/RegressionTaoAdapter.h +6 -2
  119. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/algorithms/TAO/TaoAdapter.h +6 -0
  120. sgtlearn-0.3.1/cpp/src/algorithms/TAO/TaoObjective.cpp +123 -0
  121. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/algorithms/TAO/TaoObjective.h +15 -24
  122. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/algorithms/TAO/TreeAlternatingOptimization.cpp +71 -27
  123. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/algorithms/TAO/TreeAlternatingOptimization.h +17 -9
  124. sgtlearn-0.3.1/cpp/src/algorithms/WeightedMAETree.cpp +284 -0
  125. sgtlearn-0.3.1/cpp/src/algorithms/WeightedMAETree.h +110 -0
  126. sgtlearn-0.3.1/cpp/tests/bench_mae_branch_assignment.cpp +239 -0
  127. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/tests/test_branch_assignment.cpp +40 -30
  128. sgtlearn-0.3.1/cpp/tests/test_shape_assignment_search.cpp +381 -0
  129. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/tests/test_splitters.cpp +23 -23
  130. sgtlearn-0.3.1/cpp/tests/test_weighted_mae_tree.cpp +93 -0
  131. sgtlearn-0.3.1/docs/adr/0001-regularized-outer-tree-growth.md +79 -0
  132. sgtlearn-0.3.1/docs/api/ensemble.rst +50 -0
  133. sgtlearn-0.3.1/docs/api/estimators.rst +145 -0
  134. sgtlearn-0.3.1/docs/api/plotting.rst +31 -0
  135. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/docs/api/tao.rst +20 -3
  136. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/docs/conf.py +1 -1
  137. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/docs/index.rst +7 -5
  138. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/docs/quickstart.rst +26 -2
  139. sgtlearn-0.3.1/docs/research/2026-09-09-coordinate-descent-bugfixes.md +207 -0
  140. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/docs/roadmap.rst +15 -7
  141. sgtlearn-0.3.1/docs/specs/regularized-outer-tree-growth.md +125 -0
  142. sgtlearn-0.3.1/docs/tutorials/bivariate-branching.ipynb +293 -0
  143. sgtlearn-0.3.1/docs/tutorials/categorical-features.ipynb +251 -0
  144. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/docs/tutorials/feature-importance.ipynb +31 -37
  145. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/docs/tutorials/forests.ipynb +25 -18
  146. sgtlearn-0.3.1/docs/tutorials/inspecting-trees.ipynb +230 -0
  147. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/docs/tutorials/regression.ipynb +25 -17
  148. sgtlearn-0.3.1/docs/tutorials/sgt-k.ipynb +170 -0
  149. sgtlearn-0.3.1/docs/tutorials/shape-functions.ipynb +137 -0
  150. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/docs/tutorials/structure-and-accuracy.ipynb +19 -18
  151. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/docs/tutorials/tao.ipynb +36 -28
  152. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/pyproject.toml +1 -2
  153. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/sgtlearn/_export.py +510 -31
  154. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/sgtlearn/_features.py +3 -2
  155. sgtlearn-0.3.1/sgtlearn/_multioutput.py +142 -0
  156. sgtlearn-0.3.1/sgtlearn/_weights.py +111 -0
  157. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/sgtlearn/base.py +305 -106
  158. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/sgtlearn/datasets.py +1 -3
  159. sgtlearn-0.3.1/sgtlearn/ensemble/__init__.py +4 -0
  160. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/sgtlearn/ensemble/_random_sgforest.py +93 -53
  161. sgtlearn-0.2.0/sgtlearn/ensemble/RandomSGForestClassifier.py → sgtlearn-0.3.1/sgtlearn/ensemble/random_sgforest_classifier.py +96 -35
  162. sgtlearn-0.2.0/sgtlearn/ensemble/RandomSGForestRegressor.py → sgtlearn-0.3.1/sgtlearn/ensemble/random_sgforest_regressor.py +39 -17
  163. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/sgtlearn/tao.py +62 -22
  164. sgtlearn-0.3.1/tests/discretizer_grid.py +11 -0
  165. sgtlearn-0.3.1/tests/test_branching_penalty.py +66 -0
  166. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/tests/test_categorical_classification_discretizer.py +35 -46
  167. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/tests/test_categorical_regression_discretizer.py +72 -54
  168. sgtlearn-0.3.1/tests/test_coordinate_descent_initialization.py +312 -0
  169. sgtlearn-0.3.1/tests/test_datasets.py +44 -0
  170. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/tests/test_features.py +47 -24
  171. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/tests/test_mae_regression_stress.py +14 -43
  172. sgtlearn-0.3.1/tests/test_mae_warning.py +106 -0
  173. sgtlearn-0.3.1/tests/test_missing_values.py +209 -0
  174. sgtlearn-0.3.1/tests/test_multioutput.py +77 -0
  175. sgtlearn-0.3.1/tests/test_outer_growth_integration.py +481 -0
  176. sgtlearn-0.3.1/tests/test_outer_leaf_budget.py +91 -0
  177. sgtlearn-0.3.1/tests/test_outer_penalty_validation.py +33 -0
  178. sgtlearn-0.3.1/tests/test_outer_regularized_growth.py +144 -0
  179. sgtlearn-0.3.1/tests/test_pair_screening.py +62 -0
  180. sgtlearn-0.3.1/tests/test_pairwise_categorical_missing.py +180 -0
  181. sgtlearn-0.3.1/tests/test_pairwise_classifier.py +200 -0
  182. sgtlearn-0.3.1/tests/test_pairwise_regression.py +92 -0
  183. sgtlearn-0.3.1/tests/test_plot_helpers.py +155 -0
  184. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/tests/test_plot_tree.py +128 -65
  185. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/tests/test_random_sgforest_classifier_fidelity.py +38 -13
  186. sgtlearn-0.3.1/tests/test_random_sgforest_pairwise.py +78 -0
  187. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/tests/test_random_sgforest_regressor_fidelity.py +14 -5
  188. sgtlearn-0.3.1/tests/test_random_sgforest_validation.py +132 -0
  189. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/tests/test_sgt_classifier_fidelity.py +33 -8
  190. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/tests/test_sgt_regressor_fidelity.py +16 -6
  191. sgtlearn-0.3.1/tests/test_tao.py +412 -0
  192. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/tests/test_tree_export.py +66 -5
  193. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/tests/test_univariate_classification_discretizer.py +45 -36
  194. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/tests/test_univariate_regression_discretizer.py +77 -76
  195. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/tests/test_weighted_sample.py +119 -22
  196. sgtlearn-0.2.0/cpp/src/BranchAssignmentObjectives/AbsoluteErrorBranchAssignment.cpp +0 -136
  197. sgtlearn-0.2.0/cpp/src/BranchAssignmentObjectives/AbsoluteErrorBranchAssignment.h +0 -53
  198. sgtlearn-0.2.0/cpp/src/BranchAssignmentObjectives/BranchAssignmentFactory.cpp +0 -52
  199. sgtlearn-0.2.0/cpp/src/BranchAssignmentObjectives/BranchAssignmentFactory.h +0 -32
  200. sgtlearn-0.2.0/cpp/src/Criterion.cpp +0 -124
  201. sgtlearn-0.2.0/cpp/src/Criterion.h +0 -36
  202. sgtlearn-0.2.0/cpp/src/Discretizers/ClassificationDiscretizer.h +0 -23
  203. sgtlearn-0.2.0/cpp/src/Estimators/ClassificationShapeGeneralizedTree.cpp +0 -397
  204. sgtlearn-0.2.0/cpp/src/Estimators/ClassificationShapeGeneralizedTree.h +0 -144
  205. sgtlearn-0.2.0/cpp/src/Estimators/RegressionShapeGeneralizedTree.cpp +0 -472
  206. sgtlearn-0.2.0/cpp/src/Estimators/ShapeFunctions/ShapeFunctionSplitSearch.cpp +0 -262
  207. sgtlearn-0.2.0/cpp/src/Estimators/ShapeFunctions/ShapeFunctionSplitSearch.h +0 -98
  208. sgtlearn-0.2.0/cpp/src/Splitters/categorical/CategoricalClassificationSplitter.cpp +0 -54
  209. sgtlearn-0.2.0/cpp/src/Splitters/categorical/CategoricalClassificationSplitter.h +0 -45
  210. sgtlearn-0.2.0/cpp/src/Splitters/univariate/AbsoluteErrorSplitter.cpp +0 -59
  211. sgtlearn-0.2.0/cpp/src/Splitters/univariate/AbsoluteErrorSplitter.h +0 -31
  212. sgtlearn-0.2.0/cpp/src/Splitters/univariate/ClassificationSplitter.h +0 -58
  213. sgtlearn-0.2.0/cpp/src/Splitters/univariate/SquaredErrorSplitter.cpp +0 -55
  214. sgtlearn-0.2.0/cpp/src/Splitters/univariate/SquaredErrorSplitter.h +0 -30
  215. sgtlearn-0.2.0/cpp/src/algorithms/TAO/ClassificationTaoAdapter.cpp +0 -125
  216. sgtlearn-0.2.0/cpp/src/algorithms/TAO/TaoObjective.cpp +0 -127
  217. sgtlearn-0.2.0/docs/api/ensemble.rst +0 -33
  218. sgtlearn-0.2.0/docs/api/estimators.rst +0 -55
  219. sgtlearn-0.2.0/docs/api/plotting.rst +0 -21
  220. sgtlearn-0.2.0/docs/tutorials/categorical-features.ipynb +0 -251
  221. sgtlearn-0.2.0/docs/tutorials/inspecting-trees.ipynb +0 -222
  222. sgtlearn-0.2.0/docs/tutorials/sgt-k.ipynb +0 -170
  223. sgtlearn-0.2.0/docs/tutorials/shape-functions.ipynb +0 -137
  224. sgtlearn-0.2.0/sgtlearn/_weights.py +0 -64
  225. sgtlearn-0.2.0/sgtlearn/ensemble/__init__.py +0 -4
  226. sgtlearn-0.2.0/tests/discretizer_grid.py +0 -15
  227. sgtlearn-0.2.0/tests/test_missing_values.py +0 -446
  228. sgtlearn-0.2.0/tests/test_plot_helpers.py +0 -545
  229. sgtlearn-0.2.0/tests/test_tao.py +0 -474
  230. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/.github/workflows/workflow.yml +0 -0
  231. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/.gitignore +0 -0
  232. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/.readthedocs.yaml +0 -0
  233. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/LICENSE +0 -0
  234. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/assets/SGT_Viz.png +0 -0
  235. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/README.md +0 -0
  236. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/BranchAssignmentObjectives/BranchAssignment.h +0 -0
  237. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Discretizers/DiscretizerInputKind.h +0 -0
  238. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Discretizers/univariate/GainHessianUnivariateDiscretizer.cpp +0 -0
  239. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Discretizers/univariate/GainHessianUnivariateDiscretizer.h +0 -0
  240. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Domain/CategoricalSplitCandidate.h +0 -0
  241. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Domain/FeatureInfo.h +0 -0
  242. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Domain/LearningCriterion.h +0 -0
  243. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Domain/LearningFactories.h +0 -0
  244. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Domain/UnivariateSplitCandidate.h +0 -0
  245. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Estimators/ShapeFunctions/ShapeFunctionNode.cpp +0 -0
  246. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Estimators/ShapeGeneralizedTree.h +0 -0
  247. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Splitters/categorical/CategoricalSplitter.h +0 -0
  248. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Splitters/categorical/CategoricalSplitter.tpp +0 -0
  249. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Splitters/factories/SplitterFactory.h +0 -0
  250. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Splitters/univariate/GainHessianSplitter.cpp +0 -0
  251. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Splitters/univariate/GainHessianSplitter.h +0 -0
  252. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Splitters/univariate/Splitter.tpp +0 -0
  253. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/algorithms/FeatureBagging.h +0 -0
  254. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/algorithms/ShapeBranchingTypes.h +0 -0
  255. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/algorithms/TAO/ShapeGeneralizedTaoAdapter.cpp +0 -0
  256. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/algorithms/TAO/ShapeGeneralizedTaoAdapter.h +0 -0
  257. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/algorithms/TreeBuilder.h +0 -0
  258. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/algorithms/TreeBuilder.tpp +0 -0
  259. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/algorithms/WaveletTreeMAE.cpp +0 -0
  260. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/algorithms/WaveletTreeMAE.h +0 -0
  261. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/algorithms/frontiers.h +0 -0
  262. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/algorithms/missing_values.h +0 -0
  263. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/tests/test_wavelet_tree_mae.cpp +0 -0
  264. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/docs/Makefile +0 -0
  265. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/docs/_static/.gitkeep +0 -0
  266. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/docs/_templates/.gitkeep +0 -0
  267. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/docs/api/datasets.rst +0 -0
  268. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/docs/api/index.rst +0 -0
  269. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/docs/installation.rst +0 -0
  270. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/docs/make.bat +0 -0
  271. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/docs/requirements.txt +0 -0
  272. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/example.py +0 -0
  273. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/examples/plot_tree_demo.py +0 -0
  274. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/setup.py +0 -0
  275. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/sgtlearn/__init__.py +8 -8
  276. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/sgtlearn/py.typed +0 -0
  277. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/tests/__init__.py +0 -0
  278. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/tests/constants.py +0 -0
  279. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/tests/test_max_features_float.py +0 -0
  280. {sgtlearn-0.2.0 → sgtlearn-0.3.1}/tree.png +0 -0
@@ -0,0 +1,46 @@
1
+ # Shape Generalized Trees
2
+
3
+ Language for learning trees whose branching decisions are represented by shape functions.
4
+
5
+ ## Language
6
+
7
+ **Outer tree**:
8
+ The prediction tree whose internal nodes route samples through learned shape functions.
9
+
10
+ **Inner tree**:
11
+ The tree used to partition a feature, or a pair of features, into bins for a shape function.
12
+
13
+ **Bin**:
14
+ A region produced by an inner tree whose samples share one branch assignment.
15
+
16
+ **Branch assignment**:
17
+ The mapping from inner bins to the children of an outer node.
18
+
19
+ **Branching factor**:
20
+ The upper bound on the number of children produced by an outer split.
21
+
22
+ **Split arity**:
23
+ The number of children actually produced by a particular outer split.
24
+
25
+ **Sample mass**:
26
+ The sum of sample weights in a dataset or region; it equals the sample count when every sample has weight one.
27
+
28
+ **Outer impurity**:
29
+ The arithmetic mean of per-target weighted impurity measures at an outer node. Multiplying it by sample mass gives the node's total weighted impurity contribution.
30
+
31
+ **Leaf budget**:
32
+ The maximum permitted number of outer leaves. Replacing one leaf with a split of arity k consumes k - 1 additional leaves.
33
+ _Avoid_: Node budget, when referring to `max_leaf_nodes`.
34
+
35
+ **Best-first outer growth**:
36
+ Outer-tree growth that selects the available split with the greatest positive regularized impurity improvement.
37
+ _Avoid_: BFS, which commonly means breadth-first search and does not describe this ordering.
38
+
39
+ **Regularized split improvement**:
40
+ The decrease in total weighted leaf impurity minus the additional complexity costs incurred by a split. Zero or negative improvement does not justify growth.
41
+
42
+ **Split candidate**:
43
+ A proposed outer split with its shape function, branch assignments, actual arity, and impurity improvement.
44
+
45
+ **Pair-screening proxy**:
46
+ A feasible univariate partition with positive raw impurity improvement used to estimate the promise of a feature interaction. Eligibility is independent of complexity penalties.
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.2
2
2
  Name: sgtlearn
3
- Version: 0.2.0
3
+ Version: 0.3.1
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
 
@@ -73,6 +72,22 @@ plt.show()
73
72
  ```
74
73
 
75
74
  Read the full docs here: https://sgtlearn.readthedocs.io/en/latest/index.html
75
+
76
+ Outer growth now always selects the best available regularized split and
77
+ strictly respects `max_leaf_nodes`, including multiway splits. Its impurity
78
+ improvement uses total sample weight and the mean across target outputs.
79
+ `min_impurity_decrease`, `branching_penalty` (new, default `0.0`), and
80
+ `pairwise_penalty` subtract constant costs in these units; existing positive
81
+ penalties may need retuning. Inner CART and TAO retain their separate settings.
82
+ See [outer-growth semantics](docs/api/estimators.rst) for the scoring formula.
83
+
84
+ For `SGTRegressor` and `RandomSGForestRegressor`, MAE (`criterion="absolute_error"`
85
+ or `"mae"`) leaves coordinate descent disabled by default. Each top-level fit
86
+ emits one `UserWarning`, including parallel forest fits. Set the environment
87
+ variable `SGTLEARN_MAE_CD=1` before fitting to enable it and omit the warning.
88
+ The native flag also accepts exactly `true`, `TRUE`, or `yes`; other values keep
89
+ CD disabled. This does not change TAO settings or affect classification/MSE.
90
+
76
91
  ## Developer Setup
77
92
 
78
93
  Use a **project-local virtual environment** (`.venv`) so Python, pytest, and
@@ -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
 
@@ -46,6 +45,22 @@ plt.show()
46
45
  ```
47
46
 
48
47
  Read the full docs here: https://sgtlearn.readthedocs.io/en/latest/index.html
48
+
49
+ Outer growth now always selects the best available regularized split and
50
+ strictly respects `max_leaf_nodes`, including multiway splits. Its impurity
51
+ improvement uses total sample weight and the mean across target outputs.
52
+ `min_impurity_decrease`, `branching_penalty` (new, default `0.0`), and
53
+ `pairwise_penalty` subtract constant costs in these units; existing positive
54
+ penalties may need retuning. Inner CART and TAO retain their separate settings.
55
+ See [outer-growth semantics](docs/api/estimators.rst) for the scoring formula.
56
+
57
+ For `SGTRegressor` and `RandomSGForestRegressor`, MAE (`criterion="absolute_error"`
58
+ or `"mae"`) leaves coordinate descent disabled by default. Each top-level fit
59
+ emits one `UserWarning`, including parallel forest fits. Set the environment
60
+ variable `SGTLEARN_MAE_CD=1` before fitting to enable it and omit the warning.
61
+ The native flag also accepts exactly `true`, `TRUE`, or `yes`; other values keep
62
+ CD disabled. This does not change TAO settings or affect classification/MSE.
63
+
49
64
  ## Developer Setup
50
65
 
51
66
  Use a **project-local virtual environment** (`.venv`) so Python, pytest, and
@@ -17,6 +17,18 @@
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
+ - [ ] Cross-Feature Tree Binning (Similar to [DPDT](https://github.com/KohlerHECTOR/DPDTreeEstimator)
33
+ - [ ] Boosting
34
+ - [ ] Default Optuna Support
@@ -0,0 +1,23 @@
1
+ """Runnable checks for benchmark scoring and comparison (no training required)."""
2
+ import numpy as np
3
+
4
+ from outer_growth import quality, summarize
5
+ from compare_outer_growth import performance_ratios, quality_change
6
+
7
+ y = np.array([[0, 0], [1, 1], [1, 0]])
8
+ p = np.array([[0, 1], [0, 1], [1, 0]])
9
+ assert quality(y, p, np.array([1, 2, 1]), "gini") == {
10
+ "metric": "accuracy", "mean": 0.625, "per_output": [0.5, 0.75]
11
+ }
12
+ assert quality(y, p, None, "squared_error")["mean"] == 1 / 3
13
+ assert summarize([1, 2, 3, 4, 5]) == {"median": 3.0, "mad": 1.0, "min": 1.0, "max": 5.0}
14
+ old = {k: {"median": 100} for k in ("fit_seconds", "predict_seconds_per_row", "peak_fit_process_tree_rss_bytes")}
15
+ new = {"fit_seconds": {"median": 126}, "predict_seconds_per_row": {"median": 125}, "peak_fit_process_tree_rss_bytes": {"median": 151}}
16
+ ratios = performance_ratios(old, new)
17
+ assert ratios["fit_seconds"] == {"ratio": 1.26, "investigate": True}
18
+ assert ratios["predict_seconds_per_row"] == {"ratio": 1.25, "investigate": False}
19
+ assert ratios["peak_fit_process_tree_rss_bytes"] == {"ratio": 1.51, "investigate": True}
20
+ assert quality_change({"metric": "accuracy", "mean": 0.8}, {"mean": 0.81}) == "+1.000 pp"
21
+ assert quality_change({"metric": "mse", "mean": 2}, {"mean": 1.5}) == "-0.500000 (-25.00%)"
22
+ assert quality_change({"metric": "mae", "mean": 0}, {"mean": 0.1}) == "+0.100000 (relative undefined (baseline 0))"
23
+ print("benchmark checks passed")
@@ -0,0 +1,155 @@
1
+ """Interleave matched old/new workers from outer_growth.py; no training at import."""
2
+ from __future__ import annotations
3
+
4
+ import argparse
5
+ import hashlib
6
+ import json
7
+ import os
8
+ from pathlib import Path
9
+ import subprocess
10
+
11
+ from outer_growth import report
12
+
13
+
14
+ def performance_ratios(old, new):
15
+ return {key: {"ratio": new[key]["median"] / old[key]["median"],
16
+ "investigate": new[key]["median"] / old[key]["median"] > limit}
17
+ for key, limit in (("fit_seconds", 1.25), ("predict_seconds_per_row", 1.25),
18
+ ("peak_fit_process_tree_rss_bytes", 1.5))}
19
+
20
+
21
+ def quality_change(before, after):
22
+ delta = after["mean"] - before["mean"]
23
+ if before["metric"] == "accuracy":
24
+ return f"{100 * delta:+.3f} pp"
25
+ relative = f"{100 * delta / before['mean']:+.2f}%" if before["mean"] != 0 else "relative undefined (baseline 0)"
26
+ return f"{delta:+.6f} ({relative})"
27
+
28
+
29
+ def comparison(rows):
30
+ summaries = {label: report([r for r in rows if r["build"] == label])
31
+ for label in ("baseline", "candidate")}
32
+ result = {}
33
+ for name in summaries["baseline"].keys() & summaries["candidate"].keys():
34
+ old, new = (summaries[label][name] for label in ("baseline", "candidate"))
35
+ if "fit_seconds" not in old or "fit_seconds" not in new:
36
+ continue
37
+ quality = []
38
+ for before in old["quality"]:
39
+ after = next((q for q in new["quality"] if q["seed"] == before["seed"]), None)
40
+ if after is None:
41
+ continue
42
+ delta = after["sgt"]["mean"] - before["sgt"]["mean"]
43
+ quality.append({"seed": before["seed"], "baseline": before, "candidate": after,
44
+ "candidate_minus_baseline": delta,
45
+ "training_quality_decreased": delta < 0 if before["sgt"]["metric"] == "accuracy" else delta > 0})
46
+ result[name] = {"performance": performance_ratios(old, new), "baseline": old,
47
+ "candidate": new, "quality": quality}
48
+ return dict(sorted(result.items()))
49
+
50
+
51
+ def write_report(output, rows, settings):
52
+ result = comparison(rows)
53
+ (output / "comparison.json").write_text(json.dumps(result, indent=2) + "\n")
54
+ lines = ["# Matched outer-growth comparison", "",
55
+ "Training quality is the primary soft signal. Changes below are candidate minus baseline: accuracy changes use percentage points (pp); regression losses show absolute and relative changes. Higher accuracy and lower loss are better. No numeric quality cutoff is imposed. Every worse-than-CART seed remains flagged.", ""]
56
+ if settings["candidate_branching_penalty"] is not None:
57
+ lines += [f"Separate regularization experiment: candidate branching_penalty={settings['candidate_branching_penalty']}; baseline has no public branching penalty. This is not the common zero-penalty comparison.", ""]
58
+ lines += ["| Workload / seed | Metric | Baseline | Candidate | Change | CART | Worse than CART (old/new) | Structural leaves (old/new) |",
59
+ "|---|---|---:|---:|---:|---:|---|---|"]
60
+ for name, entry in result.items():
61
+ for q in entry["quality"]:
62
+ before, after = q["baseline"], q["candidate"]
63
+ cart = f"{before['cart']['mean']:.6f}" if "cart" in before else "Excluded: grouped categorical/missing"
64
+ flags = f"{before.get('worse_than_cart', 'excluded')}/{after.get('worse_than_cart', 'excluded')}"
65
+ leaves = [[t["structural_leaves"] for t in row["structure"]] for row in (before, after)]
66
+ lines.append(f"| {name} / {q['seed']} | {before['sgt']['metric']} | {before['sgt']['mean']:.6f} | {after['sgt']['mean']:.6f} | {quality_change(before['sgt'], after['sgt'])} | {cart} | {flags} | {leaves[0]}/{leaves[1]} |")
67
+ lines += ["", "Ratios are candidate / baseline medians. Investigate fit or prediction >1.25× and peak RSS >1.50×; confirm breaches with matched reruns beyond measured noise. Review changed tree sizes alongside costs.", "",
68
+ "| Workload | Fit ratio | Predict latency ratio | Peak RSS ratio | Review gates |",
69
+ "|---|---:|---:|---:|---|"]
70
+ for name, entry in result.items():
71
+ perf = entry["performance"]
72
+ flags = [key for key, value in perf.items() if value["investigate"]]
73
+ values = [perf[key]["ratio"] for key in ("fit_seconds", "predict_seconds_per_row", "peak_fit_process_tree_rss_bytes")]
74
+ lines.append(f"| {name} | {values[0]:.3f}× | {values[1]:.3f}× | {values[2]:.3f}× | {', '.join(flags) or 'No breach'} |")
75
+ lines += ["", "Raw per-output quality, occupied leaves, depth, node counts, timing dispersion and provenance are preserved in raw.jsonl and comparison.json. Each matched pair uses the same manifest, dataset hash and seed; order alternates by round. Warmup and extra quality runs are excluded from timing medians. The runner documents memory scope in outer_growth.md.", ""]
76
+ (output / "comparison.md").write_text("\n".join(lines))
77
+
78
+
79
+ def main():
80
+ parser = argparse.ArgumentParser(description=__doc__)
81
+ parser.add_argument("--baseline-python", required=True)
82
+ parser.add_argument("--baseline-checkout", required=True)
83
+ parser.add_argument("--candidate-python", required=True)
84
+ parser.add_argument("--candidate-checkout", required=True)
85
+ parser.add_argument("--output", type=Path, required=True)
86
+ parser.add_argument("--manifest", type=Path, default=Path(__file__).with_name("outer_growth_manifest.json"))
87
+ parser.add_argument("--workload")
88
+ parser.add_argument("--candidate-branching-penalty", type=float)
89
+ args = parser.parse_args()
90
+ runner = Path(__file__).with_name("outer_growth.py").resolve()
91
+ manifest = json.loads(args.manifest.read_text())
92
+ settings = {"baseline_python": str(Path(args.baseline_python).absolute()),
93
+ "baseline_checkout": str(Path(args.baseline_checkout).resolve()),
94
+ "candidate_python": str(Path(args.candidate_python).absolute()),
95
+ "candidate_checkout": str(Path(args.candidate_checkout).resolve()),
96
+ "candidate_branching_penalty": args.candidate_branching_penalty,
97
+ "manifest": manifest, "runner_sha256": hashlib.sha256(runner.read_bytes()).hexdigest()}
98
+ for label in ("baseline", "candidate"):
99
+ settings[f"{label}_revision"] = subprocess.check_output(
100
+ ["git", "-C", settings[f"{label}_checkout"], "rev-parse", "HEAD"], text=True).strip()
101
+ args.output.mkdir(parents=True, exist_ok=True)
102
+ settings_path = args.output / "settings.json"
103
+ if settings_path.exists() and json.loads(settings_path.read_text()) != settings:
104
+ raise ValueError("Output directory belongs to different settings; use a new directory")
105
+ settings_path.write_text(json.dumps(settings, indent=2) + "\n")
106
+ raw = args.output / "raw.jsonl"
107
+ rows = [json.loads(line) for line in raw.read_text().splitlines()] if raw.exists() else []
108
+ env = dict(os.environ, **manifest["environment"])
109
+ env.pop("PYTHONPATH", None)
110
+ for work in manifest["workloads"]:
111
+ if args.workload and args.workload != work["id"]:
112
+ continue
113
+ cases = [("sgt", manifest["timing_seed"], rep) for rep in range(-1, manifest["repetitions"])]
114
+ cases += [("sgt", seed, -1) for seed in manifest["quality_seeds"] if seed != manifest["timing_seed"]]
115
+ cases += [("cart", seed, -1) for seed in manifest["quality_seeds"]]
116
+ for round_index, (model, seed, rep) in enumerate(cases):
117
+ order = ("baseline", "candidate") if round_index % 2 == 0 else ("candidate", "baseline")
118
+ pair = []
119
+ for label in order:
120
+ existing = next((r for r in rows if r["build"] == label and r["workload"] == work["id"] and r["model"] == model and r["seed"] == seed and r.get("rep", -1) == rep), None)
121
+ if existing is not None:
122
+ pair.append(existing)
123
+ continue
124
+ command = [settings[f"{label}_python"], str(runner), "--worker", "--manifest", str(args.manifest.resolve()),
125
+ "--checkout", settings[f"{label}_checkout"], "--workload", work["id"], "--model", model, "--seed", str(seed), "--rep", str(rep)]
126
+ if label == "candidate" and args.candidate_branching_penalty is not None:
127
+ command += ["--branching-penalty", str(args.candidate_branching_penalty)]
128
+ run = subprocess.run(command, cwd=settings[f"{label}_checkout"], env=env, capture_output=True, text=True, timeout=300)
129
+ with (args.output / "stderr.log").open("a") as log:
130
+ log.write(f"{label} {work['id']} {model} {seed} {rep}\n{run.stderr}\n")
131
+ if run.returncode:
132
+ raise RuntimeError(f"Worker failed: {command}\n{run.stderr}\n{run.stdout}")
133
+ row = dict(json.loads(run.stdout), build=label, round=round_index)
134
+ if label == "baseline" and "provenance" in row and row["provenance"]["revision"] != manifest["baseline_revision"]:
135
+ raise ValueError("Baseline source revision does not match frozen manifest")
136
+ rows.append(row)
137
+ pair.append(row)
138
+ with raw.open("a") as handle:
139
+ handle.write(json.dumps(row) + "\n")
140
+ print(label, work["id"], model, seed, rep, row.get("fit_seconds", row.get("excluded")), flush=True)
141
+ if all("provenance" in row for row in pair):
142
+ if pair[0]["dataset_sha256"] != pair[1]["dataset_sha256"]:
143
+ raise ValueError("Matched workers produced different datasets")
144
+ for key in ("python", "platform", "machine", "numpy", "sklearn", "environment"):
145
+ if pair[0]["provenance"][key] != pair[1]["provenance"][key]:
146
+ raise ValueError(f"Matched environments differ: {key}")
147
+ pools = [[{k: v for k, v in p.items() if k != "filepath"} for p in r["provenance"]["threadpools"]] for r in pair]
148
+ if pools[0] != pools[1]:
149
+ raise ValueError("Matched native thread pools differ")
150
+ write_report(args.output, rows, settings)
151
+ print(args.output / "comparison.md")
152
+
153
+
154
+ if __name__ == "__main__":
155
+ main()
@@ -0,0 +1,102 @@
1
+ # Outer growth baseline and comparison
2
+
3
+ Issue [#55](https://github.com/optimal-uoft/sgtlearn/issues/55), baseline commit
4
+ `90fb148e8b9b4181534dd2a19a5837aa85fba506`. Run the same manifest and runner
5
+ against separately built baseline and candidate environments. The runner rejects
6
+ SGT Python/native imports outside its virtual environment and records their SHA256
7
+ hashes, source revision, resolved estimator parameters and data hashes.
8
+
9
+ ```sh
10
+ git worktree add --detach /tmp/sgtlearn-baseline-90fb148 90fb148e8b9b4181534dd2a19a5837aa85fba506
11
+ uv venv /tmp/sgtlearn-baseline-90fb148/.venv --python 3.14.5
12
+ uv pip install --python /tmp/sgtlearn-baseline-90fb148/.venv/bin/python -r benchmarks/results/outer-growth-baseline/dependencies.txt
13
+ CMAKE_BUILD_PARALLEL_LEVEL=4 uv pip install --python /tmp/sgtlearn-baseline-90fb148/.venv/bin/python --no-build-isolation --no-deps /tmp/sgtlearn-baseline-90fb148 --config-settings build-dir=/tmp/sgtlearn-baseline-90fb148/build
14
+ /tmp/sgtlearn-baseline-90fb148/.venv/bin/python benchmarks/check_outer_growth.py
15
+ /tmp/sgtlearn-baseline-90fb148/.venv/bin/python benchmarks/outer_growth.py --checkout /tmp/sgtlearn-baseline-90fb148 --output benchmarks/results/outer-growth-baseline
16
+ ```
17
+
18
+ Use an absolute runner path when invoking from another directory. An interrupted
19
+ run resumes completed rows; use a new output directory for matched reruns or a
20
+ different build/configuration. `--workload ID` selects a workload. To compare new
21
+ branching regularization separately, use `--branching-penalty VALUE` with a new
22
+ output directory. The baseline does not support that public parameter. Equal
23
+ positive old/new `min_impurity_decrease` or `pairwise_penalty` values do **not**
24
+ represent equal regularization units. The frozen common suite uses zero penalties.
25
+
26
+ The nine bounded synthetic workloads cover Gini, entropy, MSE, MAE; binary and
27
+ multiway; capped and finite-depth uncapped growth; pairs; weights; multiple
28
+ outputs; missing/grouped categorical data; serial and two-thread forests; MAE CD
29
+ off/on; and induction only versus the unchanged automatic default TAO (10 runs).
30
+ The queue workload has depth 8 and 31 leaves; the four-way/12-leaf pair workload
31
+ forces repeated feasible-arity reductions near its cap. These characterize
32
+ behavior; public exports cannot establish private queue sizes. The duplicate-X
33
+ MAE case reuses the stress pattern in `tests/test_mae_regression_stress.py`.
34
+ The existing native `cpp/tests/bench_mae_branch_assignment.cpp` remains the
35
+ standalone branch-optimizer diagnostic; no duplicate microbenchmark is added.
36
+
37
+ Each workload has one fresh-process warmup at timing seed 17, then five
38
+ fresh-process measured repetitions. Extra warmup/quality processes use seeds 29
39
+ and 43. Raw rows retain all measurements. Startup, imports, dataset creation,
40
+ quality scoring, exports and warning collection are outside fit timing. Multiple
41
+ fits per process amortize timers where needed; their durations remain in raw
42
+ data. Predictions repeat for at least 0.25 seconds on the full training matrix.
43
+ Warnings are suppressed during timing and captured by a separate fit in quality
44
+ processes. Run while no other builds/tests are active. The OS, compiler, build
45
+ flags and exact package versions are preserved alongside results.
46
+
47
+ Peak fit memory is the maximum RSS sampled every 5 ms over the process and its
48
+ descendants, including worker memory. Existing forests use joblib threads, so
49
+ workers share the process RSS. Summing RSS can double count shared pages in a
50
+ future subprocess backend. The separate OS lifetime high-water RSS ends just
51
+ after timed fits and includes interpreter/import/data overhead. Sampling can
52
+ miss brief allocations; these figures are process footprint, not isolated native
53
+ allocation counts. Native BLAS/OpenMP threads are fixed at one; forest n_jobs
54
+ is fixed independently. Medians and median absolute deviations (MAD) are reported
55
+ with min/max; gates are review thresholds, never ordinary CI timing assertions.
56
+
57
+ Classification training accuracy is sample-weighted where weights are supplied;
58
+ multioutput results are per-output accuracy plus their arithmetic mean (not
59
+ subset/exact-match accuracy). Regression reports criterion-appropriate weighted
60
+ MSE or MAE per output and their mean. Each quality seed has a scikit-learn CART
61
+ reference on the same data/weights, with the same criterion, outer maximum depth,
62
+ maximum leaves and minimum leaf samples. CART remains binary with axis-aligned
63
+ thresholds; SGT can route noncontiguous bins, group categories, use pairs, multiple
64
+ children, TAO, and forests. Thus constraints are comparable rather than model
65
+ families equivalent. CART forests are deliberately not substituted for CART.
66
+ Grouped-categorical/missing comparison is explicitly excluded because preserving
67
+ SGT's logical category/missing routing needs different preprocessing; no silent
68
+ imputation changes the task. Exported structural/occupied leaves, depth and node
69
+ counts accompany all quality measurements (per constituent for forests).
70
+
71
+ Investigate >25% median fit/prediction slowdown or >50% peak-memory increase,
72
+ confirming any breach beyond measured noise with repeated matched runs. Interpret
73
+ performance alongside learned tree size. Training quality is the primary soft
74
+ gate: similar or improved accuracy is expected, slight degradation is acceptable,
75
+ and large reproducible drops need investigation without an invented numeric
76
+ cutoff. Lower regression loss is better. Every worse-than-CART quality result is
77
+ flagged without automatically failing the change. Default TAO and future positive
78
+ regularization comparisons stay separate from induction-only zero-cost results.
79
+
80
+ For the final common-configuration comparison, build the candidate in another
81
+ isolated checkout/venv with the same dependency pins and Release settings, then
82
+ interleave matching fresh-process runs:
83
+
84
+ ```sh
85
+ .venv/bin/python benchmarks/compare_outer_growth.py \
86
+ --baseline-python /tmp/sgtlearn-baseline-90fb148/.venv/bin/python \
87
+ --baseline-checkout /tmp/sgtlearn-baseline-90fb148 \
88
+ --candidate-python /tmp/sgtlearn-candidate/.venv/bin/python \
89
+ --candidate-checkout /tmp/sgtlearn-candidate \
90
+ --output benchmarks/results/outer-growth-comparison
91
+ ```
92
+
93
+ The driver alternates old/new order each round, verifies matching data hashes,
94
+ Python/NumPy/sklearn versions, environment and native thread pools, and writes
95
+ `raw.jsonl`, `comparison.json`, `comparison.md` and settings with revision/runner
96
+ hashes. The Markdown report lists performance review gates and every quality seed
97
+ with old/new/CART scores and leaf counts. Use `--workload ID` and a fresh output
98
+ directory to confirm a suspected breach. Optional
99
+ `--candidate-branching-penalty VALUE` produces a clearly labeled separate
100
+ regularization experiment; use another output directory. Do not change/rebuild
101
+ either environment during a run. Resume support skips completed rows; if a run
102
+ is interrupted mid-pair, use a fresh directory for strict temporal matching.