Skip to content
optimal-uoftPublic

About

sgtlearn is a Python package for learning Shape Generalized Trees (SGTs). 🌳 Shape Generalized Trees (SGTs): A class of decision trees where each node applies a learnable, potentially non-linear function, on a feature to split the data. https://arxiv.org/abs/2510.19040

Resources

Stars

15 stars

Watchers

1 watching

Forks

Repository files navigation

SGTLearn

Total sgt visualization

sgtlearn is a Python package for learning Shape Generalized Trees (SGTs).

  • 🌳 Shape Generalized Trees (SGTs): A class of decision trees where each node applies a learnable, axis-aligned shape function to one or two logical features for non-linear and interpretable splits.
  • 👁 Interpretability: Each node's shape function can be visualized directly.
  • ⚡ ShapeCART Algorithm: An efficient induction method for learning SGTs from data.
  • 🔀 Extensions:
    • Shape²GT (S²GT): Bivariate shape functions for richer splits.
    • SGTK: Multi-way branching generalization.
    • Shape²CART & ShapeCARTK: Algorithms for learning S²GTs and SGTKs.

Installation

pip install sgtlearn

Wheels are published for CPython 3.11–3.14 on Linux, macOS, and Windows (x86_64 + arm64); no compiler is needed for a binary install. To build from source instead, see Developer Setup.

Quick Start

import matplotlib.pyplot as plt
from sklearn.model_selection import train_test_split
from sgtlearn import SGTClassifier, plot_tree, make_plus

X, y = make_plus(n_samples=1500, grid=3, margin=0.07, random_state=42)

X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.2, random_state=42)

model = SGTClassifier(max_depth=4, random_state=42)
model.fit(X_train, y_train)

plot_tree(model, X=X_train)
plt.show()

Read the full docs here: https://sgtlearn.readthedocs.io/en/latest/index.html

Outer growth now always selects the best available regularized split and strictly respects max_leaf_nodes, including multiway splits. Its impurity improvement uses total sample weight and the mean across target outputs. min_impurity_decrease, branching_penalty (new, default 0.0), and pairwise_penalty subtract constant costs in these units; existing positive penalties may need retuning. Inner CART and TAO retain their separate settings. See outer-growth semantics for the scoring formula.

For SGTRegressor and RandomSGForestRegressor, MAE (criterion="absolute_error" or "mae") leaves coordinate descent disabled by default. Each top-level fit emits one UserWarning, including parallel forest fits. Set the environment variable SGTLEARN_MAE_CD=1 before fitting to enable it and omit the warning. The native flag also accepts exactly true, TRUE, or yes; other values keep CD disabled. This does not change TAO settings or affect classification/MSE.

Developer Setup

Use a project-local virtual environment (.venv) so Python, pytest, and scikit-learn stay isolated and reproducible. Pick one of the paths below (uv is recommended). All require Python ≥ 3.11.

Path 1 — uv (recommended)

uv provisions a hermetic CPython and resolves the dev extras in one step:

uv sync --all-extras
source .venv/bin/activate   # Windows: .venv\Scripts\activate

Path 2 — pip + venv (editable)

python3 -m venv .venv
source .venv/bin/activate   # Windows: .venv\Scripts\activate
pip install -U pip
pip install -e ".[dev]"

The editable install builds the C++ extensions via scikit-build-core and installs the sgtlearn package plus native modules into .venv.

Path 3 — pip non-editable (into the active environment)

pip install .
pip install ".[dev]"   # dev extras (pytest, ruff, black, mypy, ...) only if needed

Anaconda users: Do not bootstrap the venv from an Anaconda Python. Anaconda ships a libstdc++.so.6 that lags the symbol versions produced by recent system compilers (gcc ≥ 13), so the install succeeds but import sgtlearn fails with ImportError: GLIBCXX_3.4.NN not found. Use a non-Anaconda Python — e.g. uv venv --python 3.12 .venv (downloads a hermetic CPython), pyenv, or your distro's python3.

Build Workflow (scikit-build-core + CMake)

pip install . drives this build path:

  1. pyproject.toml selects scikit_build_core.build as the backend.
  2. CMake is configured from cpp/CMakeLists.txt.
  3. Each file in cpp/bindings/*.cpp becomes one pybind11 module target.
  4. After each module is built, pybind11-stubgen generates a matching .pyi (best effort; a missing or failing pybind11-stubgen does not break the build).
  5. Any generated .pyi is installed in the same location as the module .so.

C++ Folder Conventions

  • cpp/src/: C++ headers and implementation for the core library (the public include directory of sgtlearn_core).
  • cpp/bindings/: pybind11 binding entrypoints; one .cpp file maps to one Python extension module.
  • cpp/tests/: C++ unit tests consumed by the cpp_tests executable target.

CMake Targets

  • sgtlearn_core (static library): shared C++ logic used by Python modules and tests.
  • <module_name> (pybind11 module, one per file in cpp/bindings/): compiled extension modules installed as top-level modules next to the sgtlearn package.
  • cpp_tests (Catch2 executable): optional C++ test target, controlled by:
    • -DSGTLEARN_BUILD_TESTS=ON (build C++ tests)
    • -DSGTLEARN_BUILD_TESTS=OFF (default for pip install; the CMake option itself defaults to ON, but pyproject.toml overrides this so wheels don't ship test binaries)

Overriding CMake options from pip

Example (build C++ tests for one install):

pip install . --config-settings=cmake.args="-DSGTLEARN_BUILD_TESTS=ON"

License

MIT License - see LICENSE for details.

Contributing

Contributions are welcome. Please feel free to submit a pull request.

Citation

If you use this package in your research, please cite:

@article{upadhya2026empowering,
  title={Empowering Decision Trees via Shape Function Branching},
  author={Upadhya, Nakul and Cohen, Eldan},
  journal={Advances in Neural Information Processing Systems},
  volume={38},
  pages={122263--122308},
  year={2026}
}

Additionally, check out our other works on our lab website.

About

sgtlearn is a Python package for learning Shape Generalized Trees (SGTs). 🌳 Shape Generalized Trees (SGTs): A class of decision trees where each node applies a learnable, potentially non-linear function, on a feature to split the data. https://arxiv.org/abs/2510.19040

Resources

Stars

15 stars

Watchers

1 watching

Forks

Releases

Packages

Contributors

Languages