bettertrees 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (51) hide show
  1. bettertrees/__init__.py +57 -0
  2. bettertrees/_data.py +115 -0
  3. bettertrees/_typing.py +20 -0
  4. bettertrees/autotune.py +193 -0
  5. bettertrees/bins.py +247 -0
  6. bettertrees/builder.py +309 -0
  7. bettertrees/estimator.py +607 -0
  8. bettertrees/experimental/__init__.py +36 -0
  9. bettertrees/experimental/_itm_kernels.py +7 -0
  10. bettertrees/experimental/_oblique_kernels.py +29 -0
  11. bettertrees/experimental/distill.py +7 -0
  12. bettertrees/experimental/interactions.py +7 -0
  13. bettertrees/experimental/itm.py +7 -0
  14. bettertrees/experimental/oblique.py +612 -0
  15. bettertrees/experimental/precision.py +60 -0
  16. bettertrees/experimental/ratios.py +7 -0
  17. bettertrees/experimental/rulefit.py +7 -0
  18. bettertrees/kernels.py +1277 -0
  19. bettertrees/lab/__init__.py +50 -0
  20. bettertrees/lab/_itm_kernels.py +199 -0
  21. bettertrees/lab/distill.py +232 -0
  22. bettertrees/lab/interactions.py +68 -0
  23. bettertrees/lab/itm.py +726 -0
  24. bettertrees/lab/multilevel.py +400 -0
  25. bettertrees/lab/ratios.py +119 -0
  26. bettertrees/lab/robust.py +402 -0
  27. bettertrees/lab/rulefit.py +245 -0
  28. bettertrees/multilevel.py +7 -0
  29. bettertrees/postprocess.py +300 -0
  30. bettertrees/py.typed +0 -0
  31. bettertrees/search.py +388 -0
  32. bettertrees/splitters.py +191 -0
  33. bettertrees/sums/__init__.py +49 -0
  34. bettertrees/sums/_common.py +195 -0
  35. bettertrees/sums/_kernels.py +847 -0
  36. bettertrees/sums/additive.py +270 -0
  37. bettertrees/sums/budget.py +111 -0
  38. bettertrees/sums/budget.pyi +49 -0
  39. bettertrees/sums/compact.py +429 -0
  40. bettertrees/sums/edit.py +432 -0
  41. bettertrees/sums/explain.py +500 -0
  42. bettertrees/sums/imported.py +249 -0
  43. bettertrees/sums/prune.py +91 -0
  44. bettertrees/sums/robust.py +7 -0
  45. bettertrees/sums/screen.py +102 -0
  46. bettertrees/sums/smalltrees.py +668 -0
  47. bettertrees-0.1.0.dist-info/METADATA +368 -0
  48. bettertrees-0.1.0.dist-info/RECORD +51 -0
  49. bettertrees-0.1.0.dist-info/WHEEL +5 -0
  50. bettertrees-0.1.0.dist-info/licenses/LICENSE +29 -0
  51. bettertrees-0.1.0.dist-info/top_level.txt +1 -0
@@ -0,0 +1,57 @@
1
+ """Small, readable sums of trees for binary classification, under a budget of cuts.
2
+
3
+ - ``BudgetClassifier(max_splits=b)``: the recommended start, no tuning: the Interleaved
4
+ Tree Model with the fixed rule of the benchmark for a budget of b cuts.
5
+ - ``InterleavedTreeClassifier`` (the Interleaved Tree Model, ITM): a logit sum of small
6
+ trees; each step adds the best cut of any tree (or a new root) by Newton gain and
7
+ re-fits all leaves. Growth from FIGS (Tan et al., 2022), re-fit from RGF.
8
+ ``FIGSClassifier`` is an alias (the same class under its older name).
9
+ - ``FastDecisionTreeClassifier`` / ``FastDecisionTreeClassifierCV``: a single tree
10
+ (exact and histogram engines), with leaf count and shrinkage chosen by CV.
11
+ - ``SumOfOptimalTrees``: a logit sum of a few Newton-optimal trees (depth 1-3).
12
+ - ``CompactTreeBooster``: shrunken boosting of optimal trees counted in distinct
13
+ cuts (identical trees merged).
14
+ - ``AdditiveTreeBooster``: a long sum of optimal depth-1/2 trees with early stopping.
15
+ - ``LightGBMRefitClassifier`` / ``from_lightgbm``: a LightGBM model imported as an
16
+ editable sum, with its leaves refitted jointly.
17
+
18
+ The sums (binary classification) expose ``rules()``, ``explain()``, ``to_dict()``,
19
+ ``predict_contributions()``, ``plot_shapes()``, ``plot_contributions()`` and the
20
+ editing API; the single tree is multiclass and has ``export_text()``.
21
+ ``bettertrees.experimental`` has no API stability guarantee; ``bettertrees.lab`` holds
22
+ research code with negative or inconclusive results (no API, no stability, not public).
23
+ """
24
+
25
+ from .autotune import FastDecisionTreeClassifierCV
26
+ from .estimator import FastDecisionTreeClassifier
27
+ from .sums import (
28
+ AdditiveTreeBooster,
29
+ BudgetClassifier,
30
+ CompactTreeBooster,
31
+ FIGSClassifier,
32
+ InterleavedTreeClassifier,
33
+ LightGBMRefitClassifier,
34
+ SumOfOptimalTrees,
35
+ from_lightgbm,
36
+ )
37
+
38
+ __version__ = "0.1.0"
39
+
40
+ __all__ = ["AdditiveTreeBooster", "BudgetClassifier",
41
+ "CompactTreeBooster", "FIGSClassifier", "FastDecisionTreeClassifier",
42
+ "FastDecisionTreeClassifierCV", "InterleavedTreeClassifier", "LightGBMRefitClassifier",
43
+ "SumOfOptimalTrees", "__version__", "from_lightgbm"]
44
+
45
+
46
+ # moved to bettertrees.lab; still importable from here, resolved on first use (kept for
47
+ # the benchmark)
48
+ _MOVED = {"BaggedFIGSClassifier": "robust", "RashomonFIGSClassifier": "robust",
49
+ "fit_multilevel_tree": "multilevel"}
50
+
51
+
52
+ def __getattr__(name):
53
+ if name in _MOVED:
54
+ from importlib import import_module
55
+
56
+ return getattr(import_module(f".lab.{_MOVED[name]}", __name__), name)
57
+ raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
bettertrees/_data.py ADDED
@@ -0,0 +1,115 @@
1
+ """Node containers and input validation; no Numba here.
2
+
3
+ ``NodeArrays`` and ``Split`` only group arrays/scalars: there is no object per
4
+ node. Validation converts X to C-contiguous float32 and drops zero-weight rows
5
+ before any binning, support count or split.
6
+ """
7
+
8
+ from typing import NamedTuple
9
+
10
+ import numpy as np
11
+ from sklearn.utils.multiclass import check_classification_targets
12
+
13
+
14
+ class NodeArrays(NamedTuple):
15
+ """Contiguous arrays; root 0, children -1 at leaves, feature -1 at leaves.
16
+
17
+ Each position identifies a node. class_weight has shape (capacity, K) and
18
+ holds per-class masses, not probabilities. No published leaf may have zero
19
+ total mass. threshold is float64; X is float32. The builder returns the
20
+ arrays trimmed to the nodes actually used.
21
+ """
22
+
23
+ left: np.ndarray
24
+ right: np.ndarray
25
+ feature: np.ndarray
26
+ threshold: np.ndarray
27
+ missing_left: np.ndarray
28
+ class_weight: np.ndarray
29
+ n_samples: np.ndarray
30
+
31
+
32
+ class Split(NamedTuple):
33
+ """Best cut: feature=-1 means there is no admissible cut.
34
+
35
+ gain is the LOCAL Gini decrease; the builder weights it by the mass
36
+ relative to the root when applying min_impurity_decrease. threshold is
37
+ always on the original scale; bin_threshold=-1 in the exact engine. n_left
38
+ counts active rows regardless of the magnitude of their positive weights.
39
+ """
40
+
41
+ feature: int
42
+ threshold: float
43
+ missing_left: bool
44
+ gain: float
45
+ n_left: int
46
+ bin_threshold: int
47
+
48
+
49
+ def validate_X(X, *, n_features=None):
50
+ """Convert a dense numeric matrix to C-contiguous float32, allowing NaN.
51
+
52
+ Rejects empty, complex, infinite (including conversion overflow) and sparse
53
+ inputs and a wrong number of columns. Does not modify the argument. Its
54
+ cost is part of the fit/predict time of the public API.
55
+ """
56
+ raw = np.asarray(X)
57
+ if raw.ndim != 2 or 0 in raw.shape or raw.dtype.kind not in "biuf":
58
+ raise ValueError("X must be a non-empty dense numeric array.")
59
+ with np.errstate(over="ignore", invalid="ignore"):
60
+ out = np.ascontiguousarray(raw, dtype=np.float32)
61
+ if np.isinf(out).any():
62
+ raise ValueError("X cannot contain infinity or values outside the float32 range.")
63
+ if n_features is not None and out.shape[1] != n_features:
64
+ raise ValueError("The number of columns differs from training.")
65
+ return out
66
+
67
+
68
+ def prepare_training_data(X, y, sample_weight=None):
69
+ """Validate training data and return (X32, y_int32, weights64, original_classes).
70
+
71
+ Accepts a one-dimensional binary/multiclass target, including strings.
72
+ classes follows np.unique; its index breaks ties at prediction. Weights
73
+ must be finite, non-negative, aligned and have a positive sum. Zero-weight
74
+ rows take no part in bins, minimum support or splits; classes_ keeps every
75
+ class observed before that exclusion. The inputs are never modified and
76
+ no holdout is used.
77
+ """
78
+ X = validate_X(X)
79
+ target = np.asarray(y)
80
+ if target.ndim != 1 or len(target) != len(X):
81
+ raise ValueError("y must have one label per row of X.")
82
+ check_classification_targets(target)
83
+ classes, encoded = np.unique(target, return_inverse=True)
84
+ weights = (np.ones(len(X), dtype=np.float64) if sample_weight is None
85
+ else np.asarray(sample_weight, dtype=np.float64))
86
+ if weights.shape != (len(X),) or not np.isfinite(weights).all():
87
+ raise ValueError("sample_weight deve ser um vetor finito alinhado a X.")
88
+ if (weights < 0).any() or not (weights > 0).any():
89
+ raise ValueError("Sample weights must be non-negative with a positive sum.")
90
+ if not np.isfinite(weights.sum()):
91
+ raise ValueError("The sum of the weights exceeds the float64 range.")
92
+ active = weights > 0
93
+ if not active.all():
94
+ X, encoded, weights = X[active], encoded[active], weights[active]
95
+ return (np.ascontiguousarray(X), np.ascontiguousarray(encoded, dtype=np.int32),
96
+ np.ascontiguousarray(weights), classes)
97
+
98
+
99
+ def allocate_nodes(capacity, n_classes):
100
+ """Allocate node arrays for an UNTRAINED tree, all nodes initially leaves.
101
+
102
+ The builder starts small, grows the arrays geometrically and fills the
103
+ masses before publishing the tree; it avoids allocating 2**max_depth up front.
104
+ """
105
+ if capacity < 1 or n_classes < 1:
106
+ raise ValueError("capacity e n_classes devem ser positivos.")
107
+ return NodeArrays(
108
+ np.full(capacity, -1, dtype=np.int32),
109
+ np.full(capacity, -1, dtype=np.int32),
110
+ np.full(capacity, -1, dtype=np.int32),
111
+ np.full(capacity, np.nan, dtype=np.float64),
112
+ np.zeros(capacity, dtype=np.bool_),
113
+ np.zeros((capacity, n_classes), dtype=np.float64),
114
+ np.zeros(capacity, dtype=np.int64),
115
+ )
bettertrees/_typing.py ADDED
@@ -0,0 +1,20 @@
1
+ """Small shared aliases for the public API; no pandas or plotting dependencies.
2
+
3
+ ArrayLike covers NumPy arrays, sequences and objects exposing __array__, including
4
+ pandas DataFrames/Series. Shapes and column semantics are validated at runtime.
5
+ """
6
+
7
+ from collections.abc import Mapping, Sequence
8
+ from typing import Any, TypeAlias, TypeVar
9
+
10
+ import numpy as np
11
+ from numpy.typing import ArrayLike as ArrayLike
12
+ from numpy.typing import NDArray
13
+
14
+ FloatArray: TypeAlias = NDArray[np.float64]
15
+ LabelArray: TypeAlias = NDArray[Any]
16
+ IndexArray: TypeAlias = NDArray[np.intp]
17
+ FeatureNames: TypeAlias = Sequence[str] | NDArray[Any]
18
+ Seed: TypeAlias = int | None
19
+ MonotoneConstraints: TypeAlias = Mapping[int | str, int]
20
+ SelfT = TypeVar("SelfT")
@@ -0,0 +1,193 @@
1
+ """Choice of capacity and shrinkage by internal cross-validation.
2
+
3
+ A best-first tree with L leaves is the prefix of the first L-1 expansions of
4
+ the tree with L_max leaves: the expansion order does not depend on the budget.
5
+ So each inner fold needs ONE fit (with L_max); every capacity and every lambda
6
+ of hierarchical shrinkage is scored on it. The final choice is refitted on the
7
+ full training set with the winning parameters.
8
+
9
+ References
10
+ ----------
11
+ Agarwal, Tan, Ronen, Singh, Yu. "Hierarchical Shrinkage: Improving the Accuracy
12
+ and Interpretability of Tree-Based Methods." ICML 2022 (the leaf shrinkage).
13
+ Friedman, Hastie, Tibshirani. "Additive logistic regression: a statistical view
14
+ of boosting." Annals of Statistics, 2000 (best-first growth by gain).
15
+ """
16
+
17
+ from __future__ import annotations
18
+
19
+ from collections.abc import Sequence
20
+ from numbers import Integral, Real
21
+ from typing import Literal
22
+
23
+ import numpy as np
24
+ from sklearn.base import BaseEstimator, ClassifierMixin
25
+ from sklearn.model_selection import StratifiedKFold
26
+ from sklearn.utils.validation import check_is_fitted
27
+
28
+ from ._data import prepare_training_data
29
+ from ._typing import ArrayLike, FeatureNames, FloatArray, IndexArray, LabelArray, Seed, SelfT
30
+ from .estimator import FastDecisionTreeClassifier
31
+ from .postprocess import expansion_steps, hierarchical_shrinkage_probabilities, prefix_leaf_ids
32
+
33
+ # Wide grid: tuned sklearn + HS often picks > 256 leaves and lambda > 200;
34
+ # stopping at 256/200 cost 0.7% log-loss in the benchmark.
35
+ # Leaves beyond n/min_samples_leaf are never reached, so small n pays nothing.
36
+ DEFAULT_LEAVES = (4, 8, 16, 32, 64, 128, 256, 512, 1024, 2048, 4096)
37
+ DEFAULT_SHRINKAGE = (1.0, 5.0, 20.0, 50.0, 200.0, 500.0, 1000.0)
38
+
39
+
40
+ def _log_loss(y, proba, weights):
41
+ picked = np.clip(proba[np.arange(len(y)), y], 1e-15, 1.0)
42
+ return float(-np.average(np.log(picked), weights=weights))
43
+
44
+
45
+ class FastDecisionTreeClassifierCV(ClassifierMixin, BaseEstimator):
46
+ """Best-first tree whose number of leaves and shrinkage are chosen by CV.
47
+
48
+ Parameters
49
+ ----------
50
+ leaves_grid : sequence of int >= 2
51
+ Candidate capacities (maximum number of leaves).
52
+ shrinkage_grid : sequence of float > 0
53
+ Candidate ``leaf_shrinkage`` values.
54
+ cv : int >= 2
55
+ Stratified inner folds; the criterion is the validation log-loss.
56
+ Each class with positive sample weight needs at least ``cv`` rows.
57
+ Zero-weight rows do not enter the inner folds; the final tree retains
58
+ all classes observed in ``y``.
59
+ min_samples_leaf, splitter, max_bins, n_jobs, random_state
60
+ Passed to ``FastDecisionTreeClassifier`` (no pruning: truncating by
61
+ prefix needs node ids in expansion order).
62
+
63
+ Attributes
64
+ ----------
65
+ best_params_ : dict
66
+ The chosen ``max_leaf_nodes`` and ``leaf_shrinkage``.
67
+ best_estimator_ : FastDecisionTreeClassifier
68
+ The tree refitted on all rows with ``best_params_``.
69
+ cv_scores_ : ndarray of shape (len(leaves_grid), len(shrinkage_grid))
70
+ Mean validation log-loss of each pair.
71
+ classes_, n_features_in_, feature_names_in_
72
+ As in ``best_estimator_``.
73
+
74
+ Examples
75
+ --------
76
+ >>> from sklearn.datasets import load_breast_cancer
77
+ >>> from bettertrees import FastDecisionTreeClassifierCV
78
+ >>> X, y = load_breast_cancer(return_X_y=True)
79
+ >>> tree = FastDecisionTreeClassifierCV(leaves_grid=(4, 8), shrinkage_grid=(1.0, 10.0))
80
+ >>> sorted(tree.fit(X, y).best_params_)
81
+ ['leaf_shrinkage', 'max_leaf_nodes']
82
+ """
83
+
84
+ classes_: LabelArray
85
+ n_features_in_: int
86
+ feature_names_in_: LabelArray
87
+ cv_scores_: FloatArray
88
+ best_params_: dict[str, int | float]
89
+ best_estimator_: FastDecisionTreeClassifier
90
+
91
+ def __init__(self, *, leaves_grid: Sequence[int] = DEFAULT_LEAVES,
92
+ shrinkage_grid: Sequence[float] = DEFAULT_SHRINKAGE, cv: int = 3,
93
+ min_samples_leaf: int = 5, splitter: Literal["hist", "exact"] = "hist",
94
+ max_bins: int = 255, n_jobs: int = 1, random_state: Seed = 0) -> None:
95
+ self.leaves_grid = leaves_grid
96
+ self.shrinkage_grid = shrinkage_grid
97
+ self.cv = cv
98
+ self.min_samples_leaf = min_samples_leaf
99
+ self.splitter = splitter
100
+ self.max_bins = max_bins
101
+ self.n_jobs = n_jobs
102
+ self.random_state = random_state
103
+
104
+ def _tree(self, **extra):
105
+ return FastDecisionTreeClassifier(
106
+ splitter=self.splitter, min_samples_leaf=self.min_samples_leaf,
107
+ max_bins=self.max_bins, n_jobs=self.n_jobs,
108
+ random_state=self.random_state, **extra)
109
+
110
+ def fit(self: SelfT, X: ArrayLike, y: ArrayLike,
111
+ sample_weight: ArrayLike | None = None) -> SelfT:
112
+ try:
113
+ leaves = list(self.leaves_grid)
114
+ except TypeError as exc:
115
+ raise ValueError("leaves_grid must contain integers >= 2.") from exc
116
+ if not leaves or any(isinstance(v, bool | np.bool_)
117
+ or not isinstance(v, Integral) or v < 2 for v in leaves):
118
+ raise ValueError("leaves_grid must contain integers >= 2.")
119
+ try:
120
+ shrinkage = list(self.shrinkage_grid)
121
+ except TypeError as exc:
122
+ raise ValueError("shrinkage_grid must contain finite values > 0.") from exc
123
+ if not shrinkage or any(isinstance(v, bool | np.bool_)
124
+ or not isinstance(v, Real)
125
+ or not np.isfinite(v) or v <= 0 for v in shrinkage):
126
+ raise ValueError("shrinkage_grid must contain finite values > 0.")
127
+ if (isinstance(self.cv, bool | np.bool_)
128
+ or not isinstance(self.cv, Integral) or self.cv < 2):
129
+ raise ValueError("cv must be an integer >= 2.")
130
+ leaves = sorted({int(v) for v in leaves})
131
+ shrinkage = sorted({float(v) for v in shrinkage})
132
+ self._tree()._validate_parameters()
133
+ X_input = X # the final fit gets the original input, so it keeps the column names
134
+ y_input = y
135
+ X, encoded, weights, classes = prepare_training_data(X, y, sample_weight)
136
+ counts = np.bincount(encoded, minlength=len(classes))
137
+ if np.any(counts[counts > 0] < self.cv):
138
+ raise ValueError(
139
+ "Each class with positive sample weight must have at least cv rows.")
140
+ y = classes[encoded]
141
+ scores = np.zeros((len(leaves), len(shrinkage)))
142
+ folds = StratifiedKFold(n_splits=int(self.cv), shuffle=True,
143
+ random_state=self.random_state)
144
+ for fit_rows, val_rows in folds.split(X, encoded):
145
+ model = self._tree(max_leaf_nodes=leaves[-1]).fit(
146
+ X[fit_rows], y[fit_rows], sample_weight=weights[fit_rows])
147
+ nodes = model.nodes_
148
+ steps = expansion_steps(nodes)
149
+ # Classes missing from the inner fold: column with probability ~0.
150
+ columns = np.searchsorted(classes, model.classes_)
151
+ X_val = model._validate_predict_X(X[val_rows])
152
+ y_val, w_val = encoded[val_rows], weights[val_rows]
153
+ node_probs = []
154
+ for lam in shrinkage:
155
+ inner = hierarchical_shrinkage_probabilities(nodes, lam)
156
+ full = np.full((len(inner), len(classes)), 1e-15)
157
+ full[:, columns] = inner
158
+ node_probs.append(full)
159
+ for i, n_leaves in enumerate(leaves):
160
+ ids = prefix_leaf_ids(X_val, nodes, steps, n_leaves)
161
+ for j, probs in enumerate(node_probs):
162
+ scores[i, j] += _log_loss(y_val, probs[ids], w_val) / int(self.cv)
163
+ i, j = np.unravel_index(int(np.argmin(scores)), scores.shape)
164
+ best_params = dict(max_leaf_nodes=leaves[i], leaf_shrinkage=shrinkage[j])
165
+ best_estimator = self._tree(**best_params).fit(
166
+ X_input, y_input, sample_weight=sample_weight)
167
+ self.best_params_, self.cv_scores_ = best_params, scores
168
+ self.best_estimator_ = best_estimator
169
+ self.classes_ = self.best_estimator_.classes_
170
+ self.n_features_in_ = self.best_estimator_.n_features_in_
171
+ # the inner tree validates names at predict time; mirror them here (sklearn contract)
172
+ if hasattr(self.best_estimator_, "feature_names_in_"):
173
+ self.feature_names_in_ = self.best_estimator_.feature_names_in_
174
+ else:
175
+ self.__dict__.pop("feature_names_in_", None)
176
+ return self
177
+
178
+ def predict_proba(self, X: ArrayLike) -> FloatArray:
179
+ check_is_fitted(self, "best_estimator_")
180
+ return self.best_estimator_.predict_proba(X)
181
+
182
+ def predict(self, X: ArrayLike) -> LabelArray:
183
+ check_is_fitted(self, "best_estimator_")
184
+ return self.best_estimator_.predict(X)
185
+
186
+ def apply(self, X: ArrayLike) -> IndexArray:
187
+ check_is_fitted(self, "best_estimator_")
188
+ return self.best_estimator_.apply(X)
189
+
190
+ def export_text(self, feature_names: FeatureNames | None = None, precision: int = 4) -> str:
191
+ """The selected tree as text (see ``FastDecisionTreeClassifier.export_text``)."""
192
+ check_is_fitted(self, "best_estimator_")
193
+ return self.best_estimator_.export_text(feature_names, precision)
bettertrees/bins.py ADDED
@@ -0,0 +1,247 @@
1
+ """Learning and applying bin edges (NaN -> bin 0).
2
+
3
+ Quantile histograms follow the approach popularized by LightGBM (Ke et al.,
4
+ NeurIPS 2017) and scikit-learn's HistGradientBoosting: at most 255 bins per
5
+ feature, missing values in a dedicated bin.
6
+
7
+ Numba kernels in this module do not call kernels from other modules: Numba's
8
+ cache is invalidated per file and does not track cross-module dependencies.
9
+ """
10
+
11
+ from concurrent.futures import ThreadPoolExecutor
12
+
13
+ import numpy as np
14
+ from numba import njit, prange
15
+
16
+
17
+ def _count_balanced_cut_positions(counts, max_bins):
18
+ """Positions i (a cut between the i-th and the (i+1)-th distinct value).
19
+
20
+ Pure quantiles waste the budget when a few values hold most of the mass:
21
+ several quantiles fall on the same value, ``unique`` merges them, and the
22
+ tail of distinct values gets few bins. Here, values whose count is at least
23
+ the average mass per free bin get their own bins (iterating, since each
24
+ one frees budget), and the rest is split by count among the runs of light
25
+ values between them. Without ties (all counts 1) this reproduces
26
+ equal-count bins, like quantiles.
27
+ """
28
+ n_values = len(counts)
29
+ heavy = np.zeros(n_values, dtype=bool)
30
+ while True:
31
+ free_bins = max_bins - int(heavy.sum())
32
+ rest = float(counts[~heavy].sum())
33
+ if free_bins <= 0 or rest <= 0:
34
+ break
35
+ new = ~heavy & (counts >= rest / free_bins)
36
+ if not new.any():
37
+ break
38
+ heavy |= new
39
+ free_bins = max(1, max_bins - int(heavy.sum()))
40
+ target = max(float(counts[~heavy].sum()) / free_bins, 1.0)
41
+ positions = set()
42
+ for i in np.flatnonzero(heavy):
43
+ if i > 0:
44
+ positions.add(int(i) - 1)
45
+ if i < n_values - 1:
46
+ positions.add(int(i))
47
+ # Runs of light values between heavy values (or at the ends).
48
+ start = 0
49
+ while start < n_values:
50
+ if heavy[start]:
51
+ start += 1
52
+ continue
53
+ stop = start
54
+ while stop < n_values and not heavy[stop]:
55
+ stop += 1
56
+ run = counts[start:stop]
57
+ run_bins = int(round(float(run.sum()) / target))
58
+ if run_bins > 1 and stop - start > 1:
59
+ cumulative = np.cumsum(run)
60
+ goals = cumulative[-1] * np.arange(1, run_bins) / run_bins
61
+ inner = np.unique(np.searchsorted(cumulative, goals, side="left"))
62
+ positions.update(int(start + i) for i in inner if start + i < stop - 1)
63
+ start = stop
64
+ positions = np.array(sorted(positions), dtype=np.int64)
65
+ # Rounding may exceed the budget: merge the pair of neighbouring bins
66
+ # with the smallest count until it fits.
67
+ while len(positions) > max_bins - 1:
68
+ bounds = np.r_[-1, positions, n_values - 1]
69
+ cumulative = np.r_[0, np.cumsum(counts)]
70
+ sizes = cumulative[bounds[1:] + 1] - cumulative[bounds[:-1] + 1]
71
+ merged = sizes[:-1] + sizes[1:]
72
+ positions = np.delete(positions, int(np.argmin(merged)))
73
+ return positions
74
+
75
+
76
+ def _fit_bin_edges_column(col, max_bins, quantiles):
77
+ """Learn the cuts of one column; a standalone function for parallelism."""
78
+ finite = col[~np.isnan(col)]
79
+ unique, counts = np.unique(finite, return_counts=True)
80
+ unique = unique.astype(np.float64)
81
+ # Bin 0 is reserved for NaN and uint8 holds at most ids 1..255.
82
+ # So keep every interval only when they fit in the budget; high
83
+ # cardinality goes through the split by count.
84
+ if len(unique) <= max_bins:
85
+ cuts = unique[:-1] / 2 + unique[1:] / 2
86
+ else:
87
+ positions = _count_balanced_cut_positions(counts, max_bins)
88
+ cuts = unique[positions] / 2 + unique[positions + 1] / 2
89
+ return np.ascontiguousarray(cuts, dtype=np.float64)
90
+
91
+
92
+ def fit_bin_edges(X, max_bins=255, n_jobs=1):
93
+ """Learn the edges ONLY on the validated training data, ignoring NaN.
94
+
95
+ Returns a tuple of increasing float64 vectors with at most max_bins-1 cuts.
96
+ Uses midpoints when there are few unique values, which preserves indicator
97
+ columns; otherwise unweighted quantiles without sampling. Constant and
98
+ all-NaN columns get an empty vector. max_bins in [2, 255] fits uint8 with
99
+ bin 0 reserved for NaN. Weights affect impurity, not the location of these
100
+ quantiles.
101
+ """
102
+ if (isinstance(max_bins, (bool, np.bool_))
103
+ or not isinstance(max_bins, (int, np.integer))
104
+ or not 2 <= max_bins <= 255):
105
+ raise ValueError("max_bins must be an integer between 2 and 255.")
106
+ if (isinstance(n_jobs, (bool, np.bool_))
107
+ or not isinstance(n_jobs, (int, np.integer)) or n_jobs < 1):
108
+ raise ValueError("n_jobs must be a positive integer.")
109
+ quantiles = np.arange(1, max_bins) / max_bins
110
+ columns = tuple(X[:, j] for j in range(X.shape[1]))
111
+ if n_jobs == 1:
112
+ return tuple(_fit_bin_edges_column(col, max_bins, quantiles)
113
+ for col in columns)
114
+ with ThreadPoolExecutor(max_workers=int(n_jobs)) as pool:
115
+ return tuple(pool.map(
116
+ lambda col: _fit_bin_edges_column(col, max_bins, quantiles),
117
+ columns))
118
+
119
+
120
+ @njit(cache=True)
121
+ def _count_binary_values(values):
122
+ """Count zeros in a 0/1 column; return -1 on any other value."""
123
+ zeros = 0
124
+ for value in values:
125
+ if value == 0:
126
+ zeros += 1
127
+ elif value != 1:
128
+ return -1
129
+ return zeros
130
+
131
+
132
+ def _fit_bin_edges_binary(X, max_bins=255):
133
+ """Learn edges with quantiles identical to the reference on 0/1 columns.
134
+
135
+ The shortcut avoids partitioning large indicator columns. Other columns
136
+ keep the original path; `fit_bin_edges` remains the production baseline.
137
+ """
138
+ if (isinstance(max_bins, (bool, np.bool_))
139
+ or not isinstance(max_bins, (int, np.integer))
140
+ or not 2 <= max_bins <= 255):
141
+ raise ValueError("max_bins must be an integer between 2 and 255.")
142
+ edges = []
143
+ quantiles = np.arange(1, max_bins) / max_bins
144
+ for col in X.T:
145
+ finite = col[~np.isnan(col)]
146
+ unique = np.unique(finite).astype(np.float64)
147
+ if len(unique) <= max_bins:
148
+ cuts = unique[:-1] / 2 + unique[1:] / 2
149
+ else:
150
+ zeros = _count_binary_values(finite)
151
+ if zeros >= 0:
152
+ if zeros == 0 or zeros == len(finite):
153
+ cuts = np.empty(0, dtype=np.float64)
154
+ else:
155
+ positions = (len(finite) - 1) * quantiles
156
+ lower = np.floor(positions).astype(np.int64)
157
+ upper = np.ceil(positions).astype(np.int64)
158
+ cuts = np.unique(np.where(upper < zeros, 0.0,
159
+ np.where(lower >= zeros, 1.0,
160
+ positions - lower)))
161
+ cuts = cuts[cuts < 1.0]
162
+ else:
163
+ finite_min = float(finite.min())
164
+ finite_max = float(finite.max())
165
+ cuts = np.unique(np.quantile(finite, quantiles))
166
+ cuts = cuts[(cuts >= finite_min) & (cuts < finite_max)]
167
+ edges.append(np.ascontiguousarray(cuts, dtype=np.float64))
168
+ return tuple(edges)
169
+
170
+
171
+ def transform_bins(X, edges):
172
+ """Apply frozen edges and return C-contiguous uint8 with the shape of X.
173
+
174
+ NaN -> 0; finite values -> 1..B. A value equal to an edge goes to the lower
175
+ bin (searchsorted side='left'), exactly like X <= threshold at prediction.
176
+ New extremes fall into the outer bins; quantiles are never recomputed.
177
+ Internal contract: X already validated and edges produced by fit_bin_edges.
178
+ """
179
+ if len(edges) != X.shape[1]:
180
+ raise ValueError("One list of edges is required per column.")
181
+ result = np.empty(X.shape, dtype=np.uint8)
182
+ for j, cuts in enumerate(edges):
183
+ result[:, j] = np.searchsorted(cuts, X[:, j], side="left") + 1
184
+ result[np.isnan(X[:, j]), j] = 0
185
+ return result
186
+
187
+
188
+ @njit(cache=True)
189
+ def _transform_bins_row_major_kernel(X, padded_edges, lengths):
190
+ """Apply side='left' row by row, respecting the C layout of X and of the output."""
191
+ result = np.empty(X.shape, dtype=np.uint8)
192
+ for i in range(X.shape[0]):
193
+ for j in range(X.shape[1]):
194
+ value = X[i, j]
195
+ if np.isnan(value):
196
+ result[i, j] = 0
197
+ else:
198
+ lo, hi = 0, lengths[j]
199
+ while lo < hi:
200
+ mid = (lo + hi) // 2
201
+ if padded_edges[j, mid] < value:
202
+ lo = mid + 1
203
+ else:
204
+ hi = mid
205
+ result[i, j] = lo + 1
206
+ return result
207
+
208
+
209
+ @njit(cache=True, parallel=True)
210
+ def _transform_bins_row_major_parallel_kernel(X, padded_edges, lengths):
211
+ """Row-major transform, parallel over rows, without changing the cuts."""
212
+ result = np.empty(X.shape, dtype=np.uint8)
213
+ for i in prange(X.shape[0]):
214
+ for j in range(X.shape[1]):
215
+ value = X[i, j]
216
+ if np.isnan(value):
217
+ result[i, j] = 0
218
+ else:
219
+ lo, hi = 0, lengths[j]
220
+ while lo < hi:
221
+ mid = (lo + hi) // 2
222
+ if padded_edges[j, mid] < value:
223
+ lo = mid + 1
224
+ else:
225
+ hi = mid
226
+ result[i, j] = lo + 1
227
+ return result
228
+
229
+
230
+ def transform_bins_row_major(X, edges, n_jobs=1):
231
+ """Row-major variant; reproduces the NumPy reference bins exactly.
232
+
233
+ Preparing the edges is part of the measured time. `transform_bins` stays
234
+ available as a comparable baseline; both require an already validated X.
235
+ """
236
+ if len(edges) != X.shape[1]:
237
+ raise ValueError("One list of edges is required per column.")
238
+ if (isinstance(n_jobs, (bool, np.bool_)) or not isinstance(n_jobs, (int, np.integer))
239
+ or n_jobs < 1):
240
+ raise ValueError("n_jobs must be a positive integer.")
241
+ lengths = np.fromiter((len(cuts) for cuts in edges), count=len(edges), dtype=np.int64)
242
+ padded = np.zeros((len(edges), int(lengths.max(initial=0))), dtype=np.float64)
243
+ for j, cuts in enumerate(edges):
244
+ padded[j, :len(cuts)] = cuts
245
+ if n_jobs == 1:
246
+ return _transform_bins_row_major_kernel(X, padded, lengths)
247
+ return _transform_bins_row_major_parallel_kernel(X, padded, lengths)