Skip to main content

📝 scikit-learn

Description​

< What is it? >​

scikit-learn is a Python machine-learning library for supervised and unsupervised learning. It provides a consistent estimator API for preprocessing, linear models, decision trees, clustering, model selection, and evaluation. Its import package is named sklearn.

< What is fit()? >​

fit() is the standard method that trains an estimator on data. It learns or computes the information the estimator needs, stores the resulting fitted attributes on the object, and returns the fitted estimator itself. In scikit-learn, fitted attributes conventionally end with an underscore, such as coef_, tree_, or components_.

Key points​

< What happens during fitting? >​

The exact work done by fit() depends on the estimator; it is not always gradient descent or repeated epochs.

Estimator typeWhat fit() typically does
Linear regressionSolves a least-squares problem to estimate coefficients
Decision treeRecursively finds splits that improve node purity
K-means or an iterative linear modelRepeats an optimization procedure until a stopping criterion is met
Preprocessor, such as StandardScalerLearns statistics such as means and standard deviations

In all cases, fitting validates the inputs and changes the estimator from an unfitted object into one that can predict or transform data.

< The standard estimator API >​

from sklearn.linear_model import LinearRegression

model = LinearRegression()
model.fit(X_train, y_train) # learns coef_ and intercept_
predictions = model.predict(X_test)
  • fit(X, y) trains a supervised estimator using feature matrix X and targets y.
  • fit(X) fits an unsupervised estimator, such as PCA or clustering, where labels are not required.
  • predict(X) uses a fitted predictor to produce outputs for new examples.
  • transform(X) applies a fitted transformation, such as feature scaling.
  • fit_transform(X) fits a transformer and immediately transforms the training data.

For most estimators, fit() returns self, so this also works:

model = LinearRegression().fit(X_train, y_train)

< Use training data only >​

Pass training data to fit(). Use validation data to select hyperparameters and untouched test data to evaluate the final model. Fitting a scaler, transformer, or model on test data leaks information and makes evaluation unrealistically optimistic.

Comparison​

< scikit-learn and deep-learning APIs >​

ToolTypical training APIWho writes the training loop?
scikit-learnmodel.fit(X_train, y_train)The estimator
PyTorchForward pass, loss, backward(), and optimizer stepsYou usually do
Keras / TensorFlowmodel.fit(X_train, y_train, epochs=10, batch_size=32)Keras

PyTorch exposes the training loop because neural-network training often needs custom control. scikit-learn and Keras provide a higher-level fit() interface for their supported workflows.

Reference​