Using Scikit-learn Pipelines. A Titanic Walkthrough
Introduction When you build a machine learning model, you rarely feed it raw data straight away. You usually clean it first by filling in missing values, scaling numbers, and turning categories into a format the model can use. If you do each of these steps by hand, it's easy to forget one, or to accidentally let information from your test data influence how you clean your training data. A**…
Building a machine learning model typically involves cleaning raw data before feeding it into the model. This can be done manually, but it is prone to human error and can inadvertently influence the model training process. To address this issue, scikit-learn provides pipelines that automate the cleaning and modeling steps. In this article, we will explore a real-world pipeline using the Titanic dataset from Kaggle, which contains information about passengers on the ill-fated voyage. The goal is to predict survival based on features such as age, sex, ticket class, and fare.
First, we import the necessary tools, including Pipeline, ColumnTransformer, SimpleImputer, StandardScaler, OneHotEncoder, and various machine learning algorithms like logistic regression, random forest, and SVM. These lines alone do not execute any code, but they make the tools available for the subsequent steps.
The Titanic dataset consists of both numeric and categorical columns. Numeric columns include Age and Fare, while categorical columns include Sex and Embarked. Each type of column requires different preprocessing techniques. For numeric columns, we use MedianImputer to fill missing values with the median value, ensuring that extreme outliers do not skew the data.
For categorical columns, we use StrategyImputer to fill missing values with the most common category, followed by OneHotEncoder to convert categories into binary columns. This process is encapsulated in small pipelines for each column type.
ColumnTransformer is then used to apply the appropriate preprocessing pipeline to the relevant columns. It efficiently routes numeric columns through one pipeline and categorical columns through another. The resulting preprocessor object handles the entire dataset with the correct preprocessing, avoiding the pitfalls of applying the same cleaning steps uniformly across all columns.
To compare different models, we use cross-validation with StratifiedKFold, which splits the data into five folds while maintaining a balanced distribution of survived and died passengers in each fold. This is crucial because the dataset has a significantly higher number of passengers who died compared to those who survived, and an imbalanced split could lead to biased model evaluations.
By building a new pipeline for each model within each fold, we ensure that the preprocessing does not inadvertently share information between training and testing sets, thus preventing data leakage.
We evaluate each model using the F1 score, a metric that balances the trade-off between true positives and false positives, which is particularly useful when one class is more prevalent than the other. The cross_val_score function trains and tests each model five times, once per fold, and calculates the mean and standard deviation of the scores. A model with a high mean F1 score and low standard deviation is both accurate and consistent across different subsets of the data.
In our case, the Support Vector Machine (SVM) model performed the best overall, achieving a score of 0.785 in one of the folds. The Random Forest and Gradient Boosting models followed closely behind. On the other hand, the KNeighbors model was less reliable, with a lower mean score and higher variability between folds, indicating less consistency in its performance.
By examining the scores across all folds, rather than relying on a single average value, we gain a more comprehensive understanding of each model's performance and its reliability on new, unseen data. This approach, which involves building small, focused preprocessing pipelines, combining them, and rigorously testing the entire pipeline, is a valuable pattern that can be applied to various datasets and problems.
Written by urgent.news from Dev.to's reporting — not their text. Machine-written — may contain errors; check the original before relying on it.