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.
- sgtlearn-0.3.1/CONTEXT.md +46 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/PKG-INFO +20 -5
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/README.md +19 -4
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/ROADMAP.md +15 -3
- sgtlearn-0.3.1/benchmarks/check_outer_growth.py +23 -0
- sgtlearn-0.3.1/benchmarks/compare_outer_growth.py +155 -0
- sgtlearn-0.3.1/benchmarks/outer_growth.md +102 -0
- sgtlearn-0.3.1/benchmarks/outer_growth.py +261 -0
- sgtlearn-0.3.1/benchmarks/outer_growth_manifest.json +199 -0
- sgtlearn-0.3.1/benchmarks/results/outer-growth-baseline/build-provenance.txt +79 -0
- sgtlearn-0.3.1/benchmarks/results/outer-growth-baseline/dependencies.txt +28 -0
- sgtlearn-0.3.1/benchmarks/results/outer-growth-baseline/manifest.json +199 -0
- sgtlearn-0.3.1/benchmarks/results/outer-growth-baseline/raw.jsonl +99 -0
- sgtlearn-0.3.1/benchmarks/results/outer-growth-baseline/runner.sha256 +1 -0
- sgtlearn-0.3.1/benchmarks/results/outer-growth-baseline/summary.json +1154 -0
- sgtlearn-0.3.1/benchmarks/results/outer-growth-baseline/summary.md +25 -0
- sgtlearn-0.3.1/benchmarks/results/outer-growth-branching-entropy/comparison.json +391 -0
- sgtlearn-0.3.1/benchmarks/results/outer-growth-branching-entropy/comparison.md +19 -0
- sgtlearn-0.3.1/benchmarks/results/outer-growth-branching-entropy/raw.jsonl +22 -0
- sgtlearn-0.3.1/benchmarks/results/outer-growth-branching-entropy/settings.json +209 -0
- sgtlearn-0.3.1/benchmarks/results/outer-growth-branching-mse/comparison.json +415 -0
- sgtlearn-0.3.1/benchmarks/results/outer-growth-branching-mse/comparison.md +19 -0
- sgtlearn-0.3.1/benchmarks/results/outer-growth-branching-mse/raw.jsonl +22 -0
- sgtlearn-0.3.1/benchmarks/results/outer-growth-branching-mse/settings.json +209 -0
- sgtlearn-0.3.1/benchmarks/results/outer-growth-candidate-build/README.md +26 -0
- sgtlearn-0.3.1/benchmarks/results/outer-growth-candidate-build/build-provenance.txt +83 -0
- sgtlearn-0.3.1/benchmarks/results/outer-growth-candidate-build/dependencies.txt +28 -0
- sgtlearn-0.3.1/benchmarks/results/outer-growth-candidate-build/imports.json +49 -0
- sgtlearn-0.3.1/benchmarks/results/outer-growth-comparison/comparison.json +4511 -0
- sgtlearn-0.3.1/benchmarks/results/outer-growth-comparison/comparison.md +49 -0
- sgtlearn-0.3.1/benchmarks/results/outer-growth-comparison/raw.jsonl +198 -0
- sgtlearn-0.3.1/benchmarks/results/outer-growth-comparison/settings.json +209 -0
- sgtlearn-0.3.1/benchmarks/results/outer-growth-confirmation/comparison.json +2979 -0
- sgtlearn-0.3.1/benchmarks/results/outer-growth-confirmation/comparison.md +33 -0
- sgtlearn-0.3.1/benchmarks/results/outer-growth-confirmation/raw.jsonl +110 -0
- sgtlearn-0.3.1/benchmarks/results/outer-growth-confirmation/settings.json +209 -0
- sgtlearn-0.3.1/benchmarks/results/outer-growth-final.md +126 -0
- sgtlearn-0.3.1/benchmarks/results/outer-growth-verification.md +105 -0
- sgtlearn-0.3.1/benchmarks/results/outer-growth-work-diagnostic/README.md +23 -0
- sgtlearn-0.3.1/benchmarks/results/outer-growth-work-diagnostic/baseline.json +1 -0
- sgtlearn-0.3.1/benchmarks/results/outer-growth-work-diagnostic/candidate.json +1 -0
- sgtlearn-0.3.1/benchmarks/results/outer-growth-work-diagnostic/diagnostic.py +57 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/CMakeLists.txt +17 -1
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/bindings/Discretizers.cpp +147 -57
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/bindings/ShapeGeneralizedTrees.cpp +45 -21
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/bindings/TreeAlternatingOptimization.cpp +23 -19
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/bindings/_arma_bridge.h +87 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/bindings/_sgt_estimators.h +244 -52
- sgtlearn-0.3.1/cpp/src/BranchAssignmentObjectives/AbsoluteErrorBranchAssignment.cpp +214 -0
- sgtlearn-0.3.1/cpp/src/BranchAssignmentObjectives/AbsoluteErrorBranchAssignment.h +61 -0
- sgtlearn-0.3.1/cpp/src/BranchAssignmentObjectives/AbsoluteErrorBranchAssignmentBst.cpp +125 -0
- sgtlearn-0.3.1/cpp/src/BranchAssignmentObjectives/AbsoluteErrorBranchAssignmentBst.h +50 -0
- sgtlearn-0.3.1/cpp/src/BranchAssignmentObjectives/AbsoluteErrorBranchAssignmentCommon.h +51 -0
- sgtlearn-0.3.1/cpp/src/BranchAssignmentObjectives/AbsoluteErrorBranchAssignmentSort.cpp +128 -0
- sgtlearn-0.3.1/cpp/src/BranchAssignmentObjectives/AbsoluteErrorBranchAssignmentSort.h +49 -0
- sgtlearn-0.3.1/cpp/src/BranchAssignmentObjectives/BranchAssignmentFactory.cpp +98 -0
- sgtlearn-0.3.1/cpp/src/BranchAssignmentObjectives/BranchAssignmentFactory.h +43 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/BranchAssignmentObjectives/BranchAssignmentVariants.h +32 -26
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/BranchAssignmentObjectives/LeafAggregateProcessor.h +22 -14
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/BranchAssignmentObjectives/LeafAggregationBranchAssignment.cpp +4 -13
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/BranchAssignmentObjectives/LeafAggregationBranchAssignment.h +33 -2
- sgtlearn-0.3.1/cpp/src/BranchAssignmentObjectives/MaeBranchConfig.h +39 -0
- sgtlearn-0.3.1/cpp/src/Criterion.cpp +159 -0
- sgtlearn-0.3.1/cpp/src/Criterion.h +52 -0
- sgtlearn-0.3.1/cpp/src/Discretizers/ClassificationDiscretizer.h +46 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Discretizers/InnerDiscretizerBase.h +43 -13
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Discretizers/RegressionDiscretizer.h +9 -3
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Discretizers/categorical/CategoricalClassificationDiscretizer.cpp +16 -10
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Discretizers/categorical/CategoricalClassificationDiscretizer.h +9 -4
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Discretizers/categorical/CategoricalDiscretizer.h +19 -3
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Discretizers/categorical/CategoricalDiscretizer.tpp +72 -11
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Discretizers/categorical/CategoricalRegressionDiscretizer.cpp +4 -4
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Discretizers/categorical/CategoricalRegressionDiscretizer.h +6 -4
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Discretizers/factories/DiscretizerFactories.cpp +6 -5
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Discretizers/factories/DiscretizerFactories.h +6 -4
- sgtlearn-0.3.1/cpp/src/Discretizers/pair/PairClassificationDiscretizer.cpp +356 -0
- sgtlearn-0.3.1/cpp/src/Discretizers/pair/PairClassificationDiscretizer.h +106 -0
- sgtlearn-0.3.1/cpp/src/Discretizers/pair/PairRegressionDiscretizer.cpp +349 -0
- sgtlearn-0.3.1/cpp/src/Discretizers/pair/PairRegressionDiscretizer.h +69 -0
- sgtlearn-0.3.1/cpp/src/Discretizers/univariate/NumericFallbackDiscretizer.cpp +233 -0
- sgtlearn-0.3.1/cpp/src/Discretizers/univariate/NumericFallbackDiscretizer.h +19 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Discretizers/univariate/UnivariateClassificationDiscretizer.cpp +23 -14
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Discretizers/univariate/UnivariateClassificationDiscretizer.h +12 -7
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Discretizers/univariate/UnivariateDiscretizer.h +13 -2
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Discretizers/univariate/UnivariateDiscretizer.tpp +3 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Discretizers/univariate/UnivariateRegressionDiscretizer.cpp +17 -14
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Discretizers/univariate/UnivariateRegressionDiscretizer.h +6 -4
- sgtlearn-0.3.1/cpp/src/Estimators/ClassificationShapeGeneralizedTree.cpp +616 -0
- sgtlearn-0.3.1/cpp/src/Estimators/ClassificationShapeGeneralizedTree.h +192 -0
- sgtlearn-0.3.1/cpp/src/Estimators/RegressionShapeGeneralizedTree.cpp +715 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Estimators/RegressionShapeGeneralizedTree.h +45 -21
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Estimators/ShapeFunctions/NanPartitionRouting.h +83 -99
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Estimators/ShapeFunctions/ShapeFunctionNode.h +33 -7
- sgtlearn-0.3.1/cpp/src/Estimators/ShapeFunctions/ShapeFunctionSplitSearch.cpp +402 -0
- sgtlearn-0.3.1/cpp/src/Estimators/ShapeFunctions/ShapeFunctionSplitSearch.h +120 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Estimators/ShapeGeneralizedTree.cpp +7 -1
- sgtlearn-0.3.1/cpp/src/Splitters/categorical/CategoricalClassificationSplitter.cpp +57 -0
- sgtlearn-0.3.1/cpp/src/Splitters/categorical/CategoricalClassificationSplitter.h +58 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Splitters/categorical/CategoricalRegressionSplitter.cpp +26 -21
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Splitters/categorical/CategoricalRegressionSplitter.h +16 -7
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Splitters/factories/SplitterFactory.cpp +7 -2
- sgtlearn-0.3.1/cpp/src/Splitters/univariate/AbsoluteErrorSplitter.cpp +77 -0
- sgtlearn-0.3.1/cpp/src/Splitters/univariate/AbsoluteErrorSplitter.h +42 -0
- sgtlearn-0.3.1/cpp/src/Splitters/univariate/ClassificationSplitter.h +90 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Splitters/univariate/EntropySplitter.h +6 -7
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Splitters/univariate/GiniSplitter.h +6 -7
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Splitters/univariate/Splitter.h +1 -1
- sgtlearn-0.3.1/cpp/src/Splitters/univariate/SquaredErrorSplitter.cpp +71 -0
- sgtlearn-0.3.1/cpp/src/Splitters/univariate/SquaredErrorSplitter.h +42 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/algorithms/BinPartitionAssignments.h +3 -10
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/algorithms/CoordinateDescent.h +27 -8
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/algorithms/KMeansUtils.h +13 -11
- sgtlearn-0.3.1/cpp/src/algorithms/OuterTreeBuilder.h +75 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/algorithms/ShapeGeneralizedTreeParams.h +4 -10
- sgtlearn-0.3.1/cpp/src/algorithms/TAO/ClassificationTaoAdapter.cpp +210 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/algorithms/TAO/ClassificationTaoAdapter.h +15 -4
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/algorithms/TAO/RegressionTaoAdapter.cpp +79 -23
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/algorithms/TAO/RegressionTaoAdapter.h +6 -2
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/algorithms/TAO/TaoAdapter.h +6 -0
- sgtlearn-0.3.1/cpp/src/algorithms/TAO/TaoObjective.cpp +123 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/algorithms/TAO/TaoObjective.h +15 -24
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/algorithms/TAO/TreeAlternatingOptimization.cpp +71 -27
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/algorithms/TAO/TreeAlternatingOptimization.h +17 -9
- sgtlearn-0.3.1/cpp/src/algorithms/WeightedMAETree.cpp +284 -0
- sgtlearn-0.3.1/cpp/src/algorithms/WeightedMAETree.h +110 -0
- sgtlearn-0.3.1/cpp/tests/bench_mae_branch_assignment.cpp +239 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/tests/test_branch_assignment.cpp +40 -30
- sgtlearn-0.3.1/cpp/tests/test_shape_assignment_search.cpp +381 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/tests/test_splitters.cpp +23 -23
- sgtlearn-0.3.1/cpp/tests/test_weighted_mae_tree.cpp +93 -0
- sgtlearn-0.3.1/docs/adr/0001-regularized-outer-tree-growth.md +79 -0
- sgtlearn-0.3.1/docs/api/ensemble.rst +50 -0
- sgtlearn-0.3.1/docs/api/estimators.rst +145 -0
- sgtlearn-0.3.1/docs/api/plotting.rst +31 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/docs/api/tao.rst +20 -3
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/docs/conf.py +1 -1
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/docs/index.rst +7 -5
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/docs/quickstart.rst +26 -2
- sgtlearn-0.3.1/docs/research/2026-09-09-coordinate-descent-bugfixes.md +207 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/docs/roadmap.rst +15 -7
- sgtlearn-0.3.1/docs/specs/regularized-outer-tree-growth.md +125 -0
- sgtlearn-0.3.1/docs/tutorials/bivariate-branching.ipynb +293 -0
- sgtlearn-0.3.1/docs/tutorials/categorical-features.ipynb +251 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/docs/tutorials/feature-importance.ipynb +31 -37
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/docs/tutorials/forests.ipynb +25 -18
- sgtlearn-0.3.1/docs/tutorials/inspecting-trees.ipynb +230 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/docs/tutorials/regression.ipynb +25 -17
- sgtlearn-0.3.1/docs/tutorials/sgt-k.ipynb +170 -0
- sgtlearn-0.3.1/docs/tutorials/shape-functions.ipynb +137 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/docs/tutorials/structure-and-accuracy.ipynb +19 -18
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/docs/tutorials/tao.ipynb +36 -28
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/pyproject.toml +1 -2
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/sgtlearn/_export.py +510 -31
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/sgtlearn/_features.py +3 -2
- sgtlearn-0.3.1/sgtlearn/_multioutput.py +142 -0
- sgtlearn-0.3.1/sgtlearn/_weights.py +111 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/sgtlearn/base.py +305 -106
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/sgtlearn/datasets.py +1 -3
- sgtlearn-0.3.1/sgtlearn/ensemble/__init__.py +4 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/sgtlearn/ensemble/_random_sgforest.py +93 -53
- sgtlearn-0.2.0/sgtlearn/ensemble/RandomSGForestClassifier.py → sgtlearn-0.3.1/sgtlearn/ensemble/random_sgforest_classifier.py +96 -35
- sgtlearn-0.2.0/sgtlearn/ensemble/RandomSGForestRegressor.py → sgtlearn-0.3.1/sgtlearn/ensemble/random_sgforest_regressor.py +39 -17
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/sgtlearn/tao.py +62 -22
- sgtlearn-0.3.1/tests/discretizer_grid.py +11 -0
- sgtlearn-0.3.1/tests/test_branching_penalty.py +66 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/tests/test_categorical_classification_discretizer.py +35 -46
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/tests/test_categorical_regression_discretizer.py +72 -54
- sgtlearn-0.3.1/tests/test_coordinate_descent_initialization.py +312 -0
- sgtlearn-0.3.1/tests/test_datasets.py +44 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/tests/test_features.py +47 -24
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/tests/test_mae_regression_stress.py +14 -43
- sgtlearn-0.3.1/tests/test_mae_warning.py +106 -0
- sgtlearn-0.3.1/tests/test_missing_values.py +209 -0
- sgtlearn-0.3.1/tests/test_multioutput.py +77 -0
- sgtlearn-0.3.1/tests/test_outer_growth_integration.py +481 -0
- sgtlearn-0.3.1/tests/test_outer_leaf_budget.py +91 -0
- sgtlearn-0.3.1/tests/test_outer_penalty_validation.py +33 -0
- sgtlearn-0.3.1/tests/test_outer_regularized_growth.py +144 -0
- sgtlearn-0.3.1/tests/test_pair_screening.py +62 -0
- sgtlearn-0.3.1/tests/test_pairwise_categorical_missing.py +180 -0
- sgtlearn-0.3.1/tests/test_pairwise_classifier.py +200 -0
- sgtlearn-0.3.1/tests/test_pairwise_regression.py +92 -0
- sgtlearn-0.3.1/tests/test_plot_helpers.py +155 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/tests/test_plot_tree.py +128 -65
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/tests/test_random_sgforest_classifier_fidelity.py +38 -13
- sgtlearn-0.3.1/tests/test_random_sgforest_pairwise.py +78 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/tests/test_random_sgforest_regressor_fidelity.py +14 -5
- sgtlearn-0.3.1/tests/test_random_sgforest_validation.py +132 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/tests/test_sgt_classifier_fidelity.py +33 -8
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/tests/test_sgt_regressor_fidelity.py +16 -6
- sgtlearn-0.3.1/tests/test_tao.py +412 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/tests/test_tree_export.py +66 -5
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/tests/test_univariate_classification_discretizer.py +45 -36
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/tests/test_univariate_regression_discretizer.py +77 -76
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/tests/test_weighted_sample.py +119 -22
- sgtlearn-0.2.0/cpp/src/BranchAssignmentObjectives/AbsoluteErrorBranchAssignment.cpp +0 -136
- sgtlearn-0.2.0/cpp/src/BranchAssignmentObjectives/AbsoluteErrorBranchAssignment.h +0 -53
- sgtlearn-0.2.0/cpp/src/BranchAssignmentObjectives/BranchAssignmentFactory.cpp +0 -52
- sgtlearn-0.2.0/cpp/src/BranchAssignmentObjectives/BranchAssignmentFactory.h +0 -32
- sgtlearn-0.2.0/cpp/src/Criterion.cpp +0 -124
- sgtlearn-0.2.0/cpp/src/Criterion.h +0 -36
- sgtlearn-0.2.0/cpp/src/Discretizers/ClassificationDiscretizer.h +0 -23
- sgtlearn-0.2.0/cpp/src/Estimators/ClassificationShapeGeneralizedTree.cpp +0 -397
- sgtlearn-0.2.0/cpp/src/Estimators/ClassificationShapeGeneralizedTree.h +0 -144
- sgtlearn-0.2.0/cpp/src/Estimators/RegressionShapeGeneralizedTree.cpp +0 -472
- sgtlearn-0.2.0/cpp/src/Estimators/ShapeFunctions/ShapeFunctionSplitSearch.cpp +0 -262
- sgtlearn-0.2.0/cpp/src/Estimators/ShapeFunctions/ShapeFunctionSplitSearch.h +0 -98
- sgtlearn-0.2.0/cpp/src/Splitters/categorical/CategoricalClassificationSplitter.cpp +0 -54
- sgtlearn-0.2.0/cpp/src/Splitters/categorical/CategoricalClassificationSplitter.h +0 -45
- sgtlearn-0.2.0/cpp/src/Splitters/univariate/AbsoluteErrorSplitter.cpp +0 -59
- sgtlearn-0.2.0/cpp/src/Splitters/univariate/AbsoluteErrorSplitter.h +0 -31
- sgtlearn-0.2.0/cpp/src/Splitters/univariate/ClassificationSplitter.h +0 -58
- sgtlearn-0.2.0/cpp/src/Splitters/univariate/SquaredErrorSplitter.cpp +0 -55
- sgtlearn-0.2.0/cpp/src/Splitters/univariate/SquaredErrorSplitter.h +0 -30
- sgtlearn-0.2.0/cpp/src/algorithms/TAO/ClassificationTaoAdapter.cpp +0 -125
- sgtlearn-0.2.0/cpp/src/algorithms/TAO/TaoObjective.cpp +0 -127
- sgtlearn-0.2.0/docs/api/ensemble.rst +0 -33
- sgtlearn-0.2.0/docs/api/estimators.rst +0 -55
- sgtlearn-0.2.0/docs/api/plotting.rst +0 -21
- sgtlearn-0.2.0/docs/tutorials/categorical-features.ipynb +0 -251
- sgtlearn-0.2.0/docs/tutorials/inspecting-trees.ipynb +0 -222
- sgtlearn-0.2.0/docs/tutorials/sgt-k.ipynb +0 -170
- sgtlearn-0.2.0/docs/tutorials/shape-functions.ipynb +0 -137
- sgtlearn-0.2.0/sgtlearn/_weights.py +0 -64
- sgtlearn-0.2.0/sgtlearn/ensemble/__init__.py +0 -4
- sgtlearn-0.2.0/tests/discretizer_grid.py +0 -15
- sgtlearn-0.2.0/tests/test_missing_values.py +0 -446
- sgtlearn-0.2.0/tests/test_plot_helpers.py +0 -545
- sgtlearn-0.2.0/tests/test_tao.py +0 -474
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/.github/workflows/workflow.yml +0 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/.gitignore +0 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/.readthedocs.yaml +0 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/LICENSE +0 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/assets/SGT_Viz.png +0 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/README.md +0 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/BranchAssignmentObjectives/BranchAssignment.h +0 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Discretizers/DiscretizerInputKind.h +0 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Discretizers/univariate/GainHessianUnivariateDiscretizer.cpp +0 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Discretizers/univariate/GainHessianUnivariateDiscretizer.h +0 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Domain/CategoricalSplitCandidate.h +0 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Domain/FeatureInfo.h +0 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Domain/LearningCriterion.h +0 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Domain/LearningFactories.h +0 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Domain/UnivariateSplitCandidate.h +0 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Estimators/ShapeFunctions/ShapeFunctionNode.cpp +0 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Estimators/ShapeGeneralizedTree.h +0 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Splitters/categorical/CategoricalSplitter.h +0 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Splitters/categorical/CategoricalSplitter.tpp +0 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Splitters/factories/SplitterFactory.h +0 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Splitters/univariate/GainHessianSplitter.cpp +0 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Splitters/univariate/GainHessianSplitter.h +0 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/Splitters/univariate/Splitter.tpp +0 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/algorithms/FeatureBagging.h +0 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/algorithms/ShapeBranchingTypes.h +0 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/algorithms/TAO/ShapeGeneralizedTaoAdapter.cpp +0 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/algorithms/TAO/ShapeGeneralizedTaoAdapter.h +0 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/algorithms/TreeBuilder.h +0 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/algorithms/TreeBuilder.tpp +0 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/algorithms/WaveletTreeMAE.cpp +0 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/algorithms/WaveletTreeMAE.h +0 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/algorithms/frontiers.h +0 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/src/algorithms/missing_values.h +0 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/cpp/tests/test_wavelet_tree_mae.cpp +0 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/docs/Makefile +0 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/docs/_static/.gitkeep +0 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/docs/_templates/.gitkeep +0 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/docs/api/datasets.rst +0 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/docs/api/index.rst +0 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/docs/installation.rst +0 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/docs/make.bat +0 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/docs/requirements.txt +0 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/example.py +0 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/examples/plot_tree_demo.py +0 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/setup.py +0 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/sgtlearn/__init__.py +8 -8
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/sgtlearn/py.typed +0 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/tests/__init__.py +0 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/tests/constants.py +0 -0
- {sgtlearn-0.2.0 → sgtlearn-0.3.1}/tests/test_max_features_float.py +0 -0
- {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.
|
|
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
|
|
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
|
|
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
|
|
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
|
|
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
|
-
- [
|
|
21
|
-
- [
|
|
22
|
-
|
|
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.
|