Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
118 changes: 116 additions & 2 deletions ngboost/api.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,9 @@
"The NGBoost library API"

# pylint: disable=too-many-arguments
from sklearn.base import BaseEstimator
import numpy as np
from sklearn.base import BaseEstimator, ClassifierMixin
from sklearn.preprocessing import LabelEncoder
from sklearn.utils import check_array

from ngboost.distns import (
Expand Down Expand Up @@ -122,7 +124,8 @@ def __setstate__(self, state_dict):
super().__setstate__(state_dict)


class NGBClassifier(NGBoost, BaseEstimator):
# pylint: disable=duplicate-code
class NGBClassifier(ClassifierMixin, NGBoost, BaseEstimator):
"""
Constructor for NGBoost classification models.

Expand Down Expand Up @@ -154,6 +157,12 @@ class NGBClassifier(NGBoost, BaseEstimator):
tol : numerical tolerance to be used in optimization
random_state : seed for reproducibility. See
https://stackoverflow.com/questions/28064634/random-state-pseudo-random-number-in-scikit-learn
validation_fraction: Proportion of training data to set
aside as validation data for early stopping.
early_stopping_rounds: The number of consecutive boosting iterations during which the
loss has to increase before the algorithm stops early.
Set to None to disable early stopping and validation.
None enables running over the full data set.
Output:
An NGBClassifier object that can be fit.
"""
Expand All @@ -173,6 +182,8 @@ def __init__(
verbose_eval=100,
tol=1e-4,
random_state=None,
validation_fraction=0.1,
early_stopping_rounds=None,
):
assert issubclass(
Dist, ClassificationDistn
Expand All @@ -190,9 +201,112 @@ def __init__(
verbose_eval,
tol,
random_state,
validation_fraction,
early_stopping_rounds,
)
self._estimator_type = "classifier"

def _fit_label_encoder(self, Y):
le = LabelEncoder().fit(Y)
self._le = le # pylint: disable=attribute-defined-outside-init
self.classes_ = le.classes_ # pylint: disable=attribute-defined-outside-init
n_classes = self.Dist.n_params + 1
if len(self.classes_) != n_classes:
raise ValueError(
"NGBClassifier Dist expects "
f"{n_classes} classes, got {len(self.classes_)}."
)

def _encode_labels(self, Y):
return self._le.transform(Y)

def __setstate__(
self, state_dict
): # pylint: disable=attribute-defined-outside-init
super().__setstate__(state_dict)
if not hasattr(self, "classes_"):
self.classes_ = np.arange(self.Dist.n_params + 1)
if not hasattr(self, "_le"):
self._le = LabelEncoder()
self._le.classes_ = self.classes_

# pylint: disable=too-many-positional-arguments,attribute-defined-outside-init
def fit(
self,
X,
Y,
X_val=None,
Y_val=None,
sample_weight=None,
val_sample_weight=None,
train_loss_monitor=None,
val_loss_monitor=None,
early_stopping_rounds=None,
):
self._fit_label_encoder(Y)
Y = self._encode_labels(Y)
if Y_val is not None:
Y_val = self._le.transform(Y_val)
self.base_models = []
self.scalings = []
self.col_idxs = []
return NGBoost.partial_fit(
self,
X,
Y,
X_val=X_val,
Y_val=Y_val,
sample_weight=sample_weight,
val_sample_weight=val_sample_weight,
train_loss_monitor=train_loss_monitor,
val_loss_monitor=val_loss_monitor,
early_stopping_rounds=early_stopping_rounds,
)

# pylint: disable=too-many-positional-arguments,attribute-defined-outside-init
def partial_fit(
self,
X,
Y,
X_val=None,
Y_val=None,
sample_weight=None,
val_sample_weight=None,
train_loss_monitor=None,
val_loss_monitor=None,
early_stopping_rounds=None,
):
if not hasattr(self, "classes_"):
self._fit_label_encoder(Y)
Y = self._encode_labels(Y)
if Y_val is not None:
Y_val = self._le.transform(Y_val)
return NGBoost.partial_fit(
self,
X,
Y,
X_val=X_val,
Y_val=Y_val,
sample_weight=sample_weight,
val_sample_weight=val_sample_weight,
train_loss_monitor=train_loss_monitor,
val_loss_monitor=val_loss_monitor,
early_stopping_rounds=early_stopping_rounds,
)

def predict(self, X, max_iter=None):
return self.classes_[super().predict(X, max_iter=max_iter)]

def staged_predict(self, X, max_iter=None):
return [
self.classes_[pred] for pred in super().staged_predict(X, max_iter=max_iter)
]

def score(self, X, Y): # for sklearn
return self.Manifold(self.pred_param(check_array(X)).T).total_score(
self._le.transform(Y)
)

def predict_proba(self, X, max_iter=None):
"""
Probability prediction of Y at the points X=x
Expand Down
1 change: 1 addition & 0 deletions ngboost/helpers.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@
_TREE_MODULE_SWAP_LOCK = threading.RLock()


# pylint: disable-next=c-extension-no-member
class _CompatTree(_sklearn_tree.Tree): # pylint: disable=too-few-public-methods
"""Transient subclass of sklearn's Tree used only during loading.

Expand Down
50 changes: 50 additions & 0 deletions tests/test_basic.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,56 @@ def predict(self, X):
return np.full(X.shape[0], self.prediction_)


def test_classifier_sets_sklearn_classes_and_encodes_labels(breast_cancer_data):
from sklearn.base import is_classifier # pylint: disable=import-outside-toplevel
from sklearn.metrics import ( # pylint: disable=import-outside-toplevel
RocCurveDisplay,
)
from sklearn.model_selection import ( # pylint: disable=import-outside-toplevel
cross_val_score,
)

x_train, x_test, y_train, y_test = breast_cancer_data
y_labels = ["malignant" if y == 0 else "benign" for y in y_train]
y_test_labels = ["malignant" if y == 0 else "benign" for y in y_test]

ngb = NGBClassifier(
Dist=Bernoulli,
n_estimators=2,
verbose=False,
random_state=0,
)

assert is_classifier(ngb)
assert isinstance(clone(ngb), NGBClassifier)

ngb.fit(x_train, y_labels)

assert list(ngb.classes_) == ["benign", "malignant"]
assert set(ngb.predict(x_test[:10])).issubset(set(ngb.classes_))
assert ngb.predict_proba(x_test[:10]).shape == (10, 2)
assert len(ngb.staged_predict(x_test[:10])) == len(ngb.base_models)
assert cross_val_score(ngb, x_train, y_labels, scoring="roc_auc", cv=3).shape == (
3,
)

display = RocCurveDisplay.from_estimator(
ngb,
x_test,
y_test_labels,
pos_label="malignant",
)
assert display.roc_auc >= 0.5


def test_classifier_rejects_label_count_mismatch(breast_cancer_data):
x_train, _, y_train, _ = breast_cancer_data
ngb = NGBClassifier(Dist=k_categorical(3), n_estimators=2, verbose=False)

with pytest.raises(ValueError, match="expects 3 classes, got 2"):
ngb.fit(x_train, y_train)


# TODO: This is non-deterministic in the model fitting
def test_classification(breast_cancer_data):
from sklearn.metrics import ( # pylint: disable=import-outside-toplevel
Expand Down
4 changes: 3 additions & 1 deletion tests/test_helpers.py
Original file line number Diff line number Diff line change
Expand Up @@ -34,7 +34,9 @@ def _old_tree_reducer(tree):

buf = io.BytesIO()
pickler = pickle.Pickler(buf)
pickler.dispatch_table = {_sklearn_tree.Tree: _old_tree_reducer}
pickler.dispatch_table = {
_sklearn_tree.Tree: _old_tree_reducer # pylint: disable=c-extension-no-member
}
pickler.dump(model)
return buf.getvalue()

Expand Down
32 changes: 27 additions & 5 deletions tests/test_pickling.py
Original file line number Diff line number Diff line change
Expand Up @@ -59,6 +59,27 @@ def test_model_save(learners_data):
assert (new_preds == preds).all()


def test_classifier_setstate_restores_missing_label_metadata(breast_cancer_data):
"""Older classifier pickles do not include the sklearn label metadata."""

X_train, _, Y_train, _ = breast_cancer_data
ngb = NGBClassifier(verbose=False, n_estimators=2)
ngb.fit(X_train, Y_train)

state = ngb.__getstate__()
state.pop("classes_", None)
state.pop("_le", None)

model = NGBClassifier()
model.__setstate__(state)

assert np.array_equal(model.classes_, np.array([0, 1]))
label_classes = model._le.classes_ # pylint: disable=protected-access
assert np.array_equal(label_classes, model.classes_)
assert np.array_equal(model.predict(X_train[:5]), ngb.predict(X_train[:5]))
assert np.allclose(model.predict_proba(X_train[:5]), ngb.predict_proba(X_train[:5]))


# ---------------------------------------------------------------------------
# Helpers for backward-compatibility test (issue #389)
# ---------------------------------------------------------------------------
Expand Down Expand Up @@ -95,8 +116,8 @@ def _old_tree_reducer(tree):

buf = io.BytesIO()
p = pickle.Pickler(buf)
p.dispatch_table = { # pylint: disable=c-extension-no-member
_sklearn_tree.Tree: _old_tree_reducer
p.dispatch_table = {
_sklearn_tree.Tree: _old_tree_reducer # pylint: disable=c-extension-no-member
}
p.dump(model)
return buf.getvalue()
Expand Down Expand Up @@ -152,8 +173,9 @@ def test_backward_compat_load(learners_data):
assert (new_preds == preds).all()
for iter_models in model.base_models:
for estimator in iter_models:
assert isinstance( # pylint: disable=c-extension-no-member
estimator.tree_, _sklearn_tree.Tree
)
tree_type = (
_sklearn_tree.Tree
) # pylint: disable=c-extension-no-member
assert isinstance(estimator.tree_, tree_type)
finally:
os.unlink(tmp_path)
Loading