Mathew K Analytics

Lesson 16 · Scikit-learn deep dive

Scikit-learn Tutorial #16: Custom Estimators

Video sixteen of the eighteen-part series: writing your own transformer or model the scikit-learn way. BaseEstimator, TransformerMixin, and the…

⬇ Download notebookOpen in Colab ↗

📓 Full notebook

Download .ipynb

Scikit-learn Deep-Dive, Video 16: Custom Estimators#

  • Video sixteen of the eighteen-part series: writing your own transformer or model the scikit-learn way.
  • BaseEstimator, TransformerMixin, and the fit/transform/predict contract.
  • Let's get into it.

Part 1: Why Write Custom Estimators#

import numpy as np
raw_ages = np.array([[25], [-5], [200], [40], [31]])
print(raw_ages.ravel())
[ 25  -5 200  40  31]

Part 2: BaseEstimator - get_params and set_params for Free#

from sklearn.base import BaseEstimator
class ClipTransformer(BaseEstimator):
    def __init__(self, lower=0, upper=120):
        self.lower = lower
        self.upper = upper
clipper = ClipTransformer(lower=0, upper=100)
print(clipper.get_params())
{'lower': 0, 'upper': 100}

Part 3: TransformerMixin - fit_transform for Free#

from sklearn.base import TransformerMixin
class ClipTransformerV2(BaseEstimator, TransformerMixin):
    def __init__(self, lower=0, upper=120):
        self.lower = lower
        self.upper = upper
    def fit(self, X, y=None):
        return self
    def transform(self, X):
        return np.clip(X, self.lower, self.upper)
print('fit_transform' in dir(ClipTransformerV2))
True

Part 4: Using the Custom Transformer#

clipper_v2 = ClipTransformerV2(lower=0, upper=120)
cleaned_ages = clipper_v2.fit_transform(raw_ages)
print(cleaned_ages.ravel())
[ 25   0 120  40  31]

Part 5: init Rules - Store Params Unchanged, No Logic#

class BadTransformer(BaseEstimator, TransformerMixin):
    def __init__(self, factor=2):
        self.factor = factor * 2
    def fit(self, X, y=None):
        return self
    def transform(self, X):
        return X * self.factor
bad = BadTransformer(factor=3)
print(bad.get_params()['factor'], bad.factor)
6 6

Part 6: Custom Transformer with Input Validation#

from sklearn.utils.validation import check_array
class ValidatedClipper(BaseEstimator, TransformerMixin):
    def __init__(self, lower=0, upper=120):
        self.lower = lower
        self.upper = upper
    def fit(self, X, y=None):
        X = check_array(X)
        return self
    def transform(self, X):
        X = check_array(X)
        return np.clip(X, self.lower, self.upper)
validated = ValidatedClipper().fit_transform(raw_ages)
print(validated.ravel())
[ 25   0 120  40  31]

Part 7: A Custom Transformer Inside a Pipeline#

from sklearn.pipeline import Pipeline
from sklearn.preprocessing import StandardScaler
custom_pipe = Pipeline([
    ('clip', ValidatedClipper(lower=0, upper=120)),
    ('scale', StandardScaler())
])
result = custom_pipe.fit_transform(raw_ages)
print(result.ravel().round(3))
[-0.448 -1.063  1.89  -0.079 -0.3  ]

Part 8: A Custom Estimator with fit/predict#

from sklearn.base import ClassifierMixin
class MajorityClassifier(BaseEstimator, ClassifierMixin):
    def fit(self, X, y):
        vals, counts = np.unique(y, return_counts=True)
        self.majority_class_ = vals[np.argmax(counts)]
        self.classes_ = vals
        return self
    def predict(self, X):
        return np.full(shape=(len(X),), fill_value=self.majority_class_)
y_demo = np.array([0, 0, 1, 0, 0])
maj = MajorityClassifier().fit(raw_ages, y_demo)
print(maj.majority_class_, round(maj.score(raw_ages, y_demo), 3))
0 0.8

Part 9: check_is_fitted and NotFittedError#

from sklearn.utils.validation import check_is_fitted
from sklearn.exceptions import NotFittedError
class SafeMajorityClassifier(MajorityClassifier):
    def predict(self, X):
        check_is_fitted(self, 'majority_class_')
        return super().predict(X)
try:
    SafeMajorityClassifier().predict(raw_ages)
except NotFittedError as exc:
    print('caught:', type(exc).__name__)
caught: NotFittedError

Part 10: A Real Pattern - a Reusable Custom Transformer Class#

class OutlierCapper(BaseEstimator, TransformerMixin):
    def __init__(self, n_std=3):
        self.n_std = n_std
    def fit(self, X, y=None):
        X = check_array(X)
        self.mean_ = X.mean(axis=0)
        self.std_ = X.std(axis=0)
        return self
    def transform(self, X):
        X = check_array(X)
        lower = self.mean_ - self.n_std * self.std_
        upper = self.mean_ + self.n_std * self.std_
        return np.clip(X, lower, upper)
capper = OutlierCapper(n_std=2).fit(raw_ages)
print(capper.transform(raw_ages).ravel().round(1))
[ 25.  -5. 200.  40.  31.]

Wrap-Up: What You Learned#

  • Project-specific logic can be written as a proper estimator, letting it drop into Pipeline and cross_val_score.
  • BaseEstimator gives get_params and set_params automatically by inspecting the init signature.
  • TransformerMixin adds a working fit_transform automatically, as long as fit and transform are both defined.
  • init must only store parameters unchanged, with no validation or transformation logic inside it.
  • check_array validates and converts input, catching shape or type problems the way built-in transformers do.
  • A properly-built custom transformer drops into a Pipeline alongside built-in steps with no special handling.
  • A custom classifier defines fit and predict; ClassifierMixin adds a working score method automatically.
  • check_is_fitted raises a proper NotFittedError if predict or transform is called before fit.
  • Combining validation with a configurable rule makes a custom transformer genuinely reusable across projects.
  • That wraps up custom estimators. Next up: Model Persistence and Diagnostics with joblib and learning curves.

Found this useful?

All lessons, notebooks and datasets here are free. If they helped you, a coffee keeps new lessons coming.