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.
pip install sgtlearnWheels 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.
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.
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.
uv provisions a hermetic CPython and resolves the dev extras in one step:
uv sync --all-extras
source .venv/bin/activate # Windows: .venv\Scripts\activatepython3 -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.
pip install .
pip install ".[dev]" # dev extras (pytest, ruff, black, mypy, ...) only if neededAnaconda users: Do not bootstrap the venv from an Anaconda Python. Anaconda ships a
libstdc++.so.6that lags the symbol versions produced by recent system compilers (gcc ≥ 13), so the install succeeds butimport sgtlearnfails withImportError: 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'spython3.
pip install . drives this build path:
pyproject.tomlselectsscikit_build_core.buildas the backend.- CMake is configured from
cpp/CMakeLists.txt. - Each file in
cpp/bindings/*.cppbecomes one pybind11 module target. - After each module is built,
pybind11-stubgengenerates a matching.pyi(best effort; a missing or failingpybind11-stubgendoes not break the build). - Any generated
.pyiis installed in the same location as the module.so.
cpp/src/: C++ headers and implementation for the core library (the public include directory ofsgtlearn_core).cpp/bindings/: pybind11 binding entrypoints; one.cppfile maps to one Python extension module.cpp/tests/: C++ unit tests consumed by thecpp_testsexecutable target.
sgtlearn_core(static library): shared C++ logic used by Python modules and tests.<module_name>(pybind11 module, one per file incpp/bindings/): compiled extension modules installed as top-level modules next to thesgtlearnpackage.cpp_tests(Catch2 executable): optional C++ test target, controlled by:-DSGTLEARN_BUILD_TESTS=ON(build C++ tests)-DSGTLEARN_BUILD_TESTS=OFF(default forpip install; the CMake option itself defaults toON, butpyproject.tomloverrides this so wheels don't ship test binaries)
Example (build C++ tests for one install):
pip install . --config-settings=cmake.args="-DSGTLEARN_BUILD_TESTS=ON"MIT License - see LICENSE for details.
Contributions are welcome. Please feel free to submit a pull request.
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.
