scikit-learn Properly

Writing your own transformer

When no built-in tool cleans your data the way you need, a small class with fit and transform makes your own step snap into any pipeline.

On this page 5
  1. Why it exists
  2. How it works
  3. A real example you have seen
  4. Remember this
  5. What to learn next

One lesson, three depths. Pick the one that fits you today — you can switch any time.

Beginner — No maths. Plain English.

A custom transformer is your own data-cleaning step, built with the standard fit and transform plug so it snaps into any pipeline.

Think of phone chargers before USB became standard. Every brand had its own pin, and no charger fit any other phone. Then everyone agreed on one plug shape, and any charger worked with any phone.

scikit-learn's fit / transform contract is that agreed plug shape. Build your own tool with that plug, and every socket in the library accepts it — pipelines, cross-validation, grid search.

Why it exists

Real data has problems no library anticipated. One salary column has a founder paying himself a fortune. A sensor column wraps around at midnight. Your cleaning rule for these is specific to your data.

The tempting fix is a loose function that patches the table before training. But loose functions live outside the pipeline. Whatever they learn from data — a cutoff, an average — they learn from all of it, which is leakage again.

Wrap the same rule in the standard plug instead. Now the learning happens inside fit, on training data only, and the rule travels with the model.

How it works

your rule: "cap absurd salaries"
        +
standard plug: fit (learn the cap)  /  transform (apply the cap)
        =
a transformer that snaps into any pipeline, like the built-in ones

The split matters more than the rule. Whatever is learned from data happens in fit and gets remembered. Whatever is done to data happens in transform, using only what was remembered.

A real example you have seen

Water-purifier cartridges. The manufacturer publishes the socket size, so third parties can build a carbon filter, a UV stage, anything — and each slides into the same body. The purifier does not care who made the stage, because the stage honours the socket.

Remember this

  • The fit / transform contract is a plug standard anyone can build to.
  • Learning goes in fit, applying goes in transform — never the other way.
  • A custom transformer inherits pipeline superpowers: cross-validation and grid search treat it like a built-in.

What to learn next

Developer — Code and libraries.

Setup

bash
pip install scikit-learn

Tested against scikit-learn 1.7.

A cap-the-outliers transformer in 20 lines

clip_outliers.py
import numpy as np
from sklearn.base import BaseEstimator, TransformerMixin

class ClipOutliers(BaseEstimator, TransformerMixin):
    """Clip every column to percentiles learned from the training data."""

    def __init__(self, low=5.0, high=95.0):
        self.low = low            # store params untouched — a rule of the API
        self.high = high

    def fit(self, X, y=None):
        X = np.asarray(X, dtype=float)
        self.low_, self.high_ = np.percentile(X, [self.low, self.high], axis=0)
        return self               # fit returns self, so calls can chain

    def transform(self, X):
        return np.clip(np.asarray(X, dtype=float), self.low_, self.high_)

# Monthly salaries in lakhs. One founder pays himself a fortune.
salaries = np.array([[0.4], [0.5], [0.6], [0.7], [0.8], [0.9], [1.0], [45.0]])

clip = ClipOutliers()
print("ceiling learned:", clip.fit(salaries).high_.round(2))
print(clip.transform(salaries).ravel())
print(clip.transform([[300.0]]).ravel())   # a new, even wilder salary
print(clip.get_params())
Output
ceiling learned: [29.6]
[ 0.435  0.5    0.6    0.7    0.8    0.9    1.    29.6  ]
[29.6]
{'high': 95.0, 'low': 5.0}

The walkthrough

__init__ stores and does nothing else. No validation, no computation, no renaming. This is the API's strictest rule, and it exists because get_params() — printed on the last line — must round-trip the constructor arguments exactly. Grid search and clone rebuild your object from those params.

The two mixins are the plug. TransformerMixin writes fit_transform for you from your own fit and transform. BaseEstimator writes get_params, set_params, and the nice printable repr. You wrote two real methods; the ecosystem wrote the rest.

The underscore attributes appear in fit. low_ and high_ follow the learned-during-fit convention. The 300-lakh salary was capped at 29.6 — the ceiling learned during fit — not at a percentile of the new data. That is the whole anti-leakage point.

It already works everywhere. Drop it into a pipeline and the double-underscore addressing finds your params:

python
from sklearn.pipeline import make_pipeline
from sklearn.preprocessing import StandardScaler

pipe = make_pipeline(ClipOutliers(), StandardScaler())
print(pipe.get_params()["clipoutliers__high"])
Output
95.0

A grid search over {"clipoutliers__high": [90, 95, 99]} now tunes your cap honestly, refit per fold.

Common mistakes

Computing in __init__. Converting low to a fraction, validating ranges, touching data — all of it breaks cloning in ways that surface later as baffling grid-search bugs. Store verbatim; validate at the top of fit.

Learning inside transform. Calling np.percentile(X, ...) in transform recomputes the cap from whatever data arrives — including test data. It runs without error and leaks. If a value depends on data statistics, it is fit's job.

Forgetting return self. Then ClipOutliers().fit(X).transform(X) crashes with AttributeError: 'NoneType' object has no attribute 'transform', and pipelines break the same way.

Writing a class when a function was enough. A stateless step — log-transform, unit conversion — learns nothing, so it needs no class. FunctionTransformer(np.log1p) wraps any plain function into the plug shape in one line. Reach for a class only when fit must remember something.

Try it yourself

Make the transformer refuse to be fitted on data containing NaN, raising a ValueError from fit. Then add a get_feature_names_out(self, input_features=None) method returning input_features, and check the pipeline's feature names still flow through a ColumnTransformer.

What to learn next

Researcher — Mathematics and papers.

The full contract, beyond the two methods

A pipeline-compatible transformer must satisfy, in addition to fit/transform:

  • Parameter transparency: get_params() returns exactly the constructor keyword arguments; set_params(**p) accepts them. BaseEstimator implements both by introspecting the __init__ signature — which is why def __init__(self, low=5.0, high=95.0) with verbatim assignment is mandatory, not stylistic.
  • Clone semantics: clone(t) must produce an unfitted copy equivalent to type(t)(**t.get_params()). Any state created outside __init__ params and fit-time underscore attributes violates this.
  • Idempotent refit: a second fit fully resets learned state.
  • n_features_in_ and, where meaningful, get_feature_names_out — the fitted-input schema, checked by pipelines when data flows at predict time.

sklearn.utils.estimator_checks.check_estimator(ClipOutliers()) runs the conformance suite. The minimal class above fails several checks — not on logic, but on input validation: the suite expects informative errors on wrong shapes, NaN policy declared via tags, dtype preservation rules. Since 1.6, the supported route is validate_data(self, X, ...) from sklearn.utils.validation inside fit and transform, plus overriding __sklearn_tags__ to declare capabilities. Passing check_estimator is the practical bar for publishing a transformer for others; for an internal pipeline step, the four bullets above are the load-bearing subset.

Statistical footnote on the example

Winsorisation — clipping at empirical quantiles rather than deleting — dates to Winsor via Tukey's robust-statistics program; see Tukey (1962), The future of data analysis. The estimator here learns the pair (q_low, q_high) per column from the training sample. Note the estimate itself has variance O(1/(n f(q)^2)) for density f at the quantile, so with tiny n the cap is noisy — a real argument for cross-validating high rather than folklore-fixing it at 95.

Where this pattern reaches its limits

  • Transformers that must see y at transform time (target encoding) need the fit_transform-differs-from-fit-then-transform pattern with internal cross-fitting; scikit-learn's TargetEncoder documents the approach.
  • Steps that change the number of rows (resampling, outlier deletion) do not fit the transform contract at all — transform must be row-aligned. The imbalanced-learn package defines a separate fit_resample verb and its own pipeline for exactly this reason; see imbalanced data.

Reference: Buitinck et al. (2013), API design for machine learning software: experiences from the scikit-learn project; and the scikit-learn "Developing scikit-learn estimators" guide, which is the normative document for the checks.

What to learn next