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