You signed in with another tab or window. Reload to refresh your session.You signed out in another tab or window. Reload to refresh your session.You switched accounts on another tab or window. Reload to refresh your session.
Dismiss alert
<metaproperty="og:description" content="Whether you are proposing an estimator for inclusion in scikit-learn, developing a separate package compatible with scikit-learn, or implementing custom components for your own projects, this chapt..." />
<metaname="description" content="Whether you are proposing an estimator for inclusion in scikit-learn, developing a separate package compatible with scikit-learn, or implementing custom components for your own projects, this chapt..." />
<liclass="toctree-l1"><aclass="reference internal" href="plotting.html">Developing with the Plotting API</a></li>
<liclass="toctree-l1 has-children"><aclass="reference internal" href="callbacks.html">Developing with the callback API</a><details><summary><spanclass="toctree-toggle" role="presentation"><iclass="fa-solid fa-chevron-down"></i></span></summary><ul>
<liclass="toctree-l2"><aclass="reference internal" href="callback_support.html">Implementing callback support in estimators</a></li>
<spanid="develop"></span><h1>Developing scikit-learn estimators<aclass="headerlink" href="#developing-scikit-learn-estimators" title="Link to this heading">#</a></h1>
<p>Whether you are proposing an estimator for inclusion in scikit-learn,
developing a separate package compatible with scikit-learn, or
implementing custom components for your own projects, this chapter
details how to develop objects that safely interact with scikit-learn
pipelines and model selection tools.</p>
<p>This section details the public API you should use and implement for a scikit-learn
compatible estimator. Inside scikit-learn itself, we experiment and use some private
tools and our goal is always to make them public once they are stable enough, so that
you can also use them in your own projects.</p>
<sectionid="apis-of-scikit-learn-objects">
<spanid="api-overview"></span><h2>APIs of scikit-learn objects<aclass="headerlink" href="#apis-of-scikit-learn-objects" title="Link to this heading">#</a></h2>
<p>There are two major types of estimators. You can think of the first group as simple
estimators, which consists of most estimators, such as
<aclass="reference internal" href="../modules/generated/sklearn.linear_model.LogisticRegression.html#sklearn.linear_model.LogisticRegression" title="sklearn.linear_model.LogisticRegression"><codeclass="xref py py-class docutils literal notranslate"><spanclass="pre">LogisticRegression</span></code></a> or
<aclass="reference internal" href="../modules/generated/sklearn.ensemble.RandomForestClassifier.html#sklearn.ensemble.RandomForestClassifier" title="sklearn.ensemble.RandomForestClassifier"><codeclass="xref py py-class docutils literal notranslate"><spanclass="pre">RandomForestClassifier</span></code></a>. And the second group are
meta-estimators, which are estimators that wrap other estimators.
<ddclass="field-odd"><p>The base object, implements a <codeclass="docutils literal notranslate"><spanclass="pre">fit</span></code> method to learn from data, either:</p>
<p>Classification algorithms usually also offer a way to quantify certainty
of a prediction, either using <codeclass="docutils literal notranslate"><spanclass="pre">decision_function</span></code> or <codeclass="docutils literal notranslate"><spanclass="pre">predict_proba</span></code>:</p>
<ddclass="field-even"><p>A model that can give a <aclass="reference external" href="https://en.wikipedia.org/wiki/Goodness_of_fit">goodness of fit</a> measure or a likelihood of
<p>Out of all the methods that an estimator implements, <codeclass="docutils literal notranslate"><spanclass="pre">fit</span></code> is usually the one you
want to implement yourself. Other methods such as <codeclass="docutils literal notranslate"><spanclass="pre">set_params</span></code>, <codeclass="docutils literal notranslate"><spanclass="pre">get_params</span></code>, etc.
are implemented in <aclass="reference internal" href="../modules/generated/sklearn.base.BaseEstimator.html#sklearn.base.BaseEstimator" title="sklearn.base.BaseEstimator"><codeclass="xref py py-class docutils literal notranslate"><spanclass="pre">BaseEstimator</span></code></a>, which you should inherit from.
You might need to inherit from more mixins, which we will explain later.</p>
<sectionid="instantiation">
<h4>Instantiation<aclass="headerlink" href="#instantiation" title="Link to this heading">#</a></h4>
<p>This concerns the creation of an object. The object’s <codeclass="docutils literal notranslate"><spanclass="pre">__init__</span></code> method might accept
constants as arguments that determine the estimator’s behavior (like the <codeclass="docutils literal notranslate"><spanclass="pre">alpha</span></code>
constant in <aclass="reference internal" href="../modules/generated/sklearn.linear_model.SGDClassifier.html#sklearn.linear_model.SGDClassifier" title="sklearn.linear_model.SGDClassifier"><codeclass="xref py py-class docutils literal notranslate"><spanclass="pre">SGDClassifier</span></code></a>). It should not, however, take
the actual training data as an argument, as this is left to the <codeclass="docutils literal notranslate"><spanclass="pre">fit()</span></code> method:</p>
<p>Ideally, the arguments accepted by <codeclass="docutils literal notranslate"><spanclass="pre">__init__</span></code> should all be keyword arguments with a
default value. In other words, a user should be able to instantiate an estimator without
passing any arguments to it. In some cases, where there are no sane defaults for an
argument, they can be left without a default value. In scikit-learn itself, we have
very few places, only in some meta-estimators, where the sub-estimator(s) argument is
a required argument.</p>
<p>Most arguments correspond to hyperparameters describing the model or the optimisation
problem the estimator tries to solve. Other parameters might define how the estimator
behaves, e.g. defining the location of a cache to store some data. These initial
arguments (or parameters) are always remembered by the estimator. Also note that they
should not be documented under the “Attributes” section, but rather under the
<p>There should be no logic, not even input validation, and the parameters should not be
changed; which also means ideally they should not be mutable objects such as lists or
dictionaries. If they’re mutable, they should be copied before being modified. The
corresponding logic should be put where the parameters are used, typically in <codeclass="docutils literal notranslate"><spanclass="pre">fit</span></code>.
<p>The reason for postponing the validation is that if <codeclass="docutils literal notranslate"><spanclass="pre">__init__</span></code> includes input
validation, then the same validation would have to be performed in <codeclass="docutils literal notranslate"><spanclass="pre">set_params</span></code>, which
is used in algorithms like <aclass="reference internal" href="../modules/generated/sklearn.model_selection.GridSearchCV.html#sklearn.model_selection.GridSearchCV" title="sklearn.model_selection.GridSearchCV"><codeclass="xref py py-class docutils literal notranslate"><spanclass="pre">GridSearchCV</span></code></a>.</p>
<p>Also it is expected that parameters with trailing <codeclass="docutils literal notranslate"><spanclass="pre">_</span></code> are <strong>not to be set
inside the</strong><codeclass="docutils literal notranslate"><spanclass="pre">__init__</span></code><strong>method</strong>. More details on attributes that are not init
arguments come shortly.</p>
</section>
<sectionid="fitting">
<h4>Fitting<aclass="headerlink" href="#fitting" title="Link to this heading">#</a></h4>
<p>The next thing you will probably want to do is to estimate some parameters in the model.
This is implemented in the <codeclass="docutils literal notranslate"><spanclass="pre">fit()</span></code> method, and it’s where the training happens.
For instance, this is where you have the computation to learn or estimate coefficients
for a linear model.</p>
<p>The <codeclass="docutils literal notranslate"><spanclass="pre">fit()</span></code> method takes the training data as arguments, which can be one
array in the case of unsupervised learning, or two arrays in the case
of supervised learning. Other metadata that come with the training data, such as
<codeclass="docutils literal notranslate"><spanclass="pre">sample_weight</span></code>, can also be passed to <codeclass="docutils literal notranslate"><spanclass="pre">fit</span></code> as keyword arguments.</p>
<p>Note that the model is fitted using <codeclass="docutils literal notranslate"><spanclass="pre">X</span></code> and <codeclass="docutils literal notranslate"><spanclass="pre">y</span></code>, but the object holds no
reference to <codeclass="docutils literal notranslate"><spanclass="pre">X</span></code> and <codeclass="docutils literal notranslate"><spanclass="pre">y</span></code>. There are, however, some exceptions to this, as in
the case of precomputed kernels where this data must be stored for use by
<p>The number of samples, i.e. <codeclass="docutils literal notranslate"><spanclass="pre">X.shape[0]</span></code> should be the same as <codeclass="docutils literal notranslate"><spanclass="pre">y.shape[0]</span></code>. If this
requirement is not met, an exception of type <codeclass="docutils literal notranslate"><spanclass="pre">ValueError</span></code> should be raised.</p>
<p><codeclass="docutils literal notranslate"><spanclass="pre">y</span></code> might be ignored in the case of unsupervised learning. However, to
make it possible to use the estimator as part of a pipeline that can
mix both supervised and unsupervised transformers, even unsupervised
estimators need to accept a <codeclass="docutils literal notranslate"><spanclass="pre">y=None</span></code> keyword argument in
the second position that is just ignored by the estimator.
For the same reason, <codeclass="docutils literal notranslate"><spanclass="pre">fit_predict</span></code>, <codeclass="docutils literal notranslate"><spanclass="pre">fit_transform</span></code>, <codeclass="docutils literal notranslate"><spanclass="pre">score</span></code>
and <codeclass="docutils literal notranslate"><spanclass="pre">partial_fit</span></code> methods need to accept a <codeclass="docutils literal notranslate"><spanclass="pre">y</span></code> argument in
the second place if they are implemented.</p>
<p>The method should return the object (<codeclass="docutils literal notranslate"><spanclass="pre">self</span></code>). This pattern is useful
to be able to implement quick one liners in an IPython session such as:</p>
<p>Depending on the nature of the algorithm, <codeclass="docutils literal notranslate"><spanclass="pre">fit</span></code> can sometimes also accept additional
keywords arguments. However, any parameter that can have a value assigned prior to
having access to the data should be an <codeclass="docutils literal notranslate"><spanclass="pre">__init__</span></code> keyword argument. Ideally, <strong>fit
parameters should be restricted to directly data dependent variables</strong>. For instance a
Gram matrix or an affinity matrix which are precomputed from the data matrix <codeclass="docutils literal notranslate"><spanclass="pre">X</span></code> are
data dependent. A tolerance stopping criterion <codeclass="docutils literal notranslate"><spanclass="pre">tol</span></code> is not directly data dependent
(although the optimal value according to some scoring function probably is).</p>
<p>When <codeclass="docutils literal notranslate"><spanclass="pre">fit</span></code> is called, any previous call to <codeclass="docutils literal notranslate"><spanclass="pre">fit</span></code> should be ignored. In
general, calling <codeclass="docutils literal notranslate"><spanclass="pre">estimator.fit(X1)</span></code> and then <codeclass="docutils literal notranslate"><spanclass="pre">estimator.fit(X2)</span></code> should
be the same as only calling <codeclass="docutils literal notranslate"><spanclass="pre">estimator.fit(X2)</span></code>. However, this may not be
true in practice when <codeclass="docutils literal notranslate"><spanclass="pre">fit</span></code> depends on some random process, see
<aclass="reference internal" href="../glossary.html#term-random_state"><spanclass="xref std std-term">random_state</span></a>. Another exception to this rule is when the
hyper-parameter <codeclass="docutils literal notranslate"><spanclass="pre">warm_start</span></code> is set to <codeclass="docutils literal notranslate"><spanclass="pre">True</span></code> for estimators that
support it. <codeclass="docutils literal notranslate"><spanclass="pre">warm_start=True</span></code> means that the previous state of the
trainable parameters of the estimator are reused instead of using the
default initialization strategy.</p>
</section>
<sectionid="estimated-attributes">
<h4>Estimated Attributes<aclass="headerlink" href="#estimated-attributes" title="Link to this heading">#</a></h4>
<p>According to scikit-learn conventions, attributes which you’d want to expose to your
users as public attributes and have been estimated or learned from the data must always
have a name ending with trailing underscore, for example the coefficients of some
regression estimator would be stored in a <codeclass="docutils literal notranslate"><spanclass="pre">coef_</span></code> attribute after <codeclass="docutils literal notranslate"><spanclass="pre">fit</span></code> has been
called. Similarly, attributes that you learn in the process and you’d like to store yet
not expose to the user, should have a leading underscore, e.g. <codeclass="docutils literal notranslate"><spanclass="pre">_intermediate_coefs</span></code>.
You’d need to document the first group (with a trailing underscore) as “Attributes” and
no need to document the second group (with a leading underscore).</p>
<p>The estimated attributes are expected to be overridden when you call <codeclass="docutils literal notranslate"><spanclass="pre">fit</span></code> a second
time.</p>
</section>
<sectionid="universal-attributes">
<h4>Universal attributes<aclass="headerlink" href="#universal-attributes" title="Link to this heading">#</a></h4>
<p>Estimators that expect tabular input should set a <codeclass="docutils literal notranslate"><spanclass="pre">n_features_in_</span></code>
attribute at <codeclass="docutils literal notranslate"><spanclass="pre">fit</span></code> time to indicate the number of features that the estimator
expects for subsequent calls to <aclass="reference internal" href="../glossary.html#term-predict"><spanclass="xref std std-term">predict</span></a> or <aclass="reference internal" href="../glossary.html#term-transform"><spanclass="xref std std-term">transform</span></a>.
See <aclass="reference external" href="https://scikit-learn-enhancement-proposals.readthedocs.io/en/latest/slep010/proposal.html">SLEP010</a>
for details.</p>
<p>Similarly, if estimators are given dataframes such as pandas or polars, they should
set a <codeclass="docutils literal notranslate"><spanclass="pre">feature_names_in_</span></code> attribute to indicate the features names of the input data,
detailed in <aclass="reference external" href="https://scikit-learn-enhancement-proposals.readthedocs.io/en/latest/slep007/proposal.html">SLEP007</a>.
Using <aclass="reference internal" href="../modules/generated/sklearn.utils.validation.validate_data.html#sklearn.utils.validation.validate_data" title="sklearn.utils.validation.validate_data"><codeclass="xref py py-func docutils literal notranslate"><spanclass="pre">validate_data</span></code></a> would automatically set these
attributes for you.</p>
</section>
</section>
</section>
<sectionid="rolling-your-own-estimator">
<spanid="id1"></span><h2>Rolling your own estimator<aclass="headerlink" href="#rolling-your-own-estimator" title="Link to this heading">#</a></h2>
<p>If you want to implement a new estimator that is scikit-learn compatible, there are
several internals of scikit-learn that you should be aware of in addition to
the scikit-learn API outlined above. You can check whether your estimator
adheres to the scikit-learn interface and standards by running
<aclass="reference internal" href="../modules/generated/sklearn.utils.estimator_checks.check_estimator.html#sklearn.utils.estimator_checks.check_estimator" title="sklearn.utils.estimator_checks.check_estimator"><codeclass="xref py py-func docutils literal notranslate"><spanclass="pre">check_estimator</span></code></a> on an instance. The
<p>The main motivation to make a class compatible to the scikit-learn estimator
interface might be that you want to use it together with model evaluation and
selection tools such as <aclass="reference internal" href="../modules/generated/sklearn.model_selection.GridSearchCV.html#sklearn.model_selection.GridSearchCV" title="sklearn.model_selection.GridSearchCV"><codeclass="xref py py-class docutils literal notranslate"><spanclass="pre">GridSearchCV</span></code></a> and
<p>Before detailing the required interface below, we describe two ways to achieve
the correct interface more easily.</p>
<asideclass="topic">
<pclass="topic-title">Project template:</p>
<p>We provide a <aclass="reference external" href="https://github.com/scikit-learn-contrib/project-template/">project template</a> which helps in the
creation of Python packages containing scikit-learn compatible estimators. It
provides:</p>
<ulclass="simple">
<li><p>an initial git repository with Python package directory structure</p></li>
<li><p>a template of a scikit-learn estimator</p></li>
<li><p>an initial test suite including use of <codeclass="xref py py-func docutils literal notranslate"><spanclass="pre">parametrize_with_checks</span></code></p></li>
<li><p>directory structures and scripts to compile documentation and example
galleries</p></li>
<li><p>scripts to manage continuous integration (testing on Linux, MacOS, and Windows)</p></li>
<li><p>instructions from getting started to publishing on <aclass="reference external" href="https://pypi.org/">PyPi</a></p></li>
<p>We tend to use “duck typing” instead of checking for <aclass="reference external" href="https://docs.python.org/3/builtins/functions.html#isinstance" title="(in Python v3.14)"><codeclass="xref py py-func docutils literal notranslate"><spanclass="pre">isinstance</span></code></a>, which means
it’s technically possible to implement an estimator without inheriting from
scikit-learn classes. However, if you don’t inherit from the right mixins, either
there will be a large amount of boilerplate code for you to implement and keep in
sync with scikit-learn development, or your estimator might not function the same
way as a scikit-learn estimator. Here we only document how to develop an estimator
using our mixins. If you’re interested in implementing your estimator without
inheriting from scikit-learn mixins, you’d need to check our implementations.</p>
<p>For example, below is a custom classifier, with more examples included in the
<h3>get_params and set_params<aclass="headerlink" href="#get-params-and-set-params" title="Link to this heading">#</a></h3>
<p>All scikit-learn estimators have <codeclass="docutils literal notranslate"><spanclass="pre">get_params</span></code> and <codeclass="docutils literal notranslate"><spanclass="pre">set_params</span></code> functions.</p>
<p>The <codeclass="docutils literal notranslate"><spanclass="pre">get_params</span></code> function takes no arguments and returns a dict of the
<codeclass="docutils literal notranslate"><spanclass="pre">__init__</span></code> parameters of the estimator, together with their values.</p>
<p>It takes one keyword argument, <codeclass="docutils literal notranslate"><spanclass="pre">deep</span></code>, which receives a boolean value that determines
whether the method should return the parameters of sub-estimators (only relevant for
meta-estimators). The default value for <codeclass="docutils literal notranslate"><spanclass="pre">deep</span></code> is <codeclass="docutils literal notranslate"><spanclass="pre">True</span></code>. For instance considering