Breast Cancer Wisconsin Diagnostic — ML Analysis
A reproducible machine learning pipeline that classifies breast masses as malignant or benign from FNA biopsy measurements, with clinical threshold analysis and SHAP interpretability.
Introduction
Dataset Overview
The Breast Cancer Wisconsin (Diagnostic) dataset contains 569 samples with 30 numerical features computed from digitized images of Fine Needle Aspiration (FNA) biopsies of breast masses. For each cell nucleus, ten real-valued features are computed (radius, texture, perimeter, area, smoothness, compactness, concavity, concave points, symmetry, fractal dimension), and three statistics are reported: mean, standard error (SE), and worst (largest of the three values), yielding 30 features in total.
The target variable is diagnosis: M (malignant, encoded as 1) or B (benign, encoded as 0).
Clinical Relevance
Breast cancer is the most commonly diagnosed cancer in women worldwide. Early and accurate detection is critical — distinguishing malignant from benign masses using FNA biopsy analysis can inform timely treatment decisions, potentially saving lives. An ML classifier on these nuclear measurements can assist pathologists by providing an objective, reproducible second opinion.
In a screening context, minimizing false negatives (missed malignancies) is paramount — it is clinically safer to recall a benign case than to miss a cancer. In a confirmatory context, precision becomes more important to avoid unnecessary invasive follow-up procedures.
Citation
Wolberg, W., Mangasarian, O., Street, N. & Street, W. (1993). Breast Cancer Wisconsin (Diagnostic). UCI Machine Learning Repository. https://doi.org/10.24432/C5DW2B
Data Loading & First Look
>_ Show codeIn [1]
import os
import sys
import warnings
import pickle
import numpy as np
import pandas as pd
import matplotlib
import matplotlib.pyplot as plt
import seaborn as sns
import shap
from sklearn.linear_model import LogisticRegression
from sklearn.tree import DecisionTreeClassifier
from sklearn.ensemble import RandomForestClassifier, GradientBoostingClassifier
from sklearn.model_selection import GridSearchCV
from sklearn.metrics import (
accuracy_score, precision_score, recall_score,
f1_score, roc_auc_score, confusion_matrix,
roc_curve, precision_recall_curve,
)
warnings.filterwarnings('ignore')
# Add src to path
SRC_DIR = os.path.join(os.path.dirname(os.path.abspath('.')), 'breast-cancer-ml-analysis-2', 'src')
BASE_DIR = '/home/gabriele/cancer-research/breast-cancer-ml-analysis-2'
sys.path.insert(0, os.path.join(BASE_DIR, 'src'))
from preprocessing import load_data, encode_target, split_features_target, get_train_test_split, fit_scaler, apply_scaler
from evaluation import compute_metrics, build_metrics_table, plot_confusion_matrix, plot_roc_curves, plot_precision_recall_curve, plot_threshold_analysis
RANDOM_STATE = 42
INPUT_PATH = os.path.join(BASE_DIR, 'input', 'wdbc.csv')
PLOTS_DIR = os.path.join(BASE_DIR, 'outputs', 'plots')
MODELS_DIR = os.path.join(BASE_DIR, 'outputs', 'models')
PLOT_STYLE = 'seaborn-v0_8-whitegrid'
plt.style.use(PLOT_STYLE)
print('Libraries loaded successfully')Libraries loaded successfully
>_ Show codeIn [2]
# Load and inspect raw data
df_raw = load_data(INPUT_PATH, id_cols=['id'])
print(f'Shape: {df_raw.shape}')
print(f'Columns: {df_raw.columns.tolist()}')
df_raw.head()Shape: (569, 31) Columns: ['diagnosis', 'radius_mean', 'texture_mean', 'perimeter_mean', 'area_mean', 'smoothness_mean', 'compactness_mean', 'concavity_mean', 'concave points_mean', 'symmetry_mean', 'fractal_dimension_mean', 'radius_se', 'texture_se', 'perimeter_se', 'area_se', 'smoothness_se', 'compactness_se', 'concavity_se', 'concave points_se', 'symmetry_se', 'fractal_dimension_se', 'radius_worst', 'texture_worst', 'perimeter_worst', 'area_worst', 'smoothness_worst', 'compactness_worst', 'concavity_worst', 'concave points_worst', 'symmetry_worst', 'fractal_dimension_worst']
| diagnosis | radius_mean | texture_mean | perimeter_mean | area_mean | smoothness_mean | compactness_mean | concavity_mean | concave points_mean | symmetry_mean | ... | radius_worst | texture_worst | perimeter_worst | area_worst | smoothness_worst | compactness_worst | concavity_worst | concave points_worst | symmetry_worst | fractal_dimension_worst | |
|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|
| 0 | M | 17.99 | 10.38 | 122.80 | 1001.0 | 0.11840 | 0.27760 | 0.3001 | 0.14710 | 0.2419 | ... | 25.38 | 17.33 | 184.60 | 2019.0 | 0.1622 | 0.6656 | 0.7119 | 0.2654 | 0.4601 | 0.11890 |
| 1 | M | 20.57 | 17.77 | 132.90 | 1326.0 | 0.08474 | 0.07864 | 0.0869 | 0.07017 | 0.1812 | ... | 24.99 | 23.41 | 158.80 | 1956.0 | 0.1238 | 0.1866 | 0.2416 | 0.1860 | 0.2750 | 0.08902 |
| 2 | M | 19.69 | 21.25 | 130.00 | 1203.0 | 0.10960 | 0.15990 | 0.1974 | 0.12790 | 0.2069 | ... | 23.57 | 25.53 | 152.50 | 1709.0 | 0.1444 | 0.4245 | 0.4504 | 0.2430 | 0.3613 | 0.08758 |
| 3 | M | 11.42 | 20.38 | 77.58 | 386.1 | 0.14250 | 0.28390 | 0.2414 | 0.10520 | 0.2597 | ... | 14.91 | 26.50 | 98.87 | 567.7 | 0.2098 | 0.8663 | 0.6869 | 0.2575 | 0.6638 | 0.17300 |
| 4 | M | 20.29 | 14.34 | 135.10 | 1297.0 | 0.10030 | 0.13280 | 0.1980 | 0.10430 | 0.1809 | ... | 22.54 | 16.67 | 152.20 | 1575.0 | 0.1374 | 0.2050 | 0.4000 | 0.1625 | 0.2364 | 0.07678 |
>_ Show codeIn [3]
# Data types and basic statistics
print('Data types:')
print(df_raw.dtypes)
print(f'\nMissing values: {df_raw.isnull().sum().sum()}')Data types: diagnosis str radius_mean float64 texture_mean float64 perimeter_mean float64 area_mean float64 smoothness_mean float64 compactness_mean float64 concavity_mean float64 concave points_mean float64 symmetry_mean float64 fractal_dimension_mean float64 radius_se float64 texture_se float64 perimeter_se float64 area_se float64 smoothness_se float64 compactness_se float64 concavity_se float64 concave points_se float64 symmetry_se float64 fractal_dimension_se float64 radius_worst float64 texture_worst float64 perimeter_worst float64 area_worst float64 smoothness_worst float64 compactness_worst float64 concavity_worst float64 concave points_worst float64 symmetry_worst float64 fractal_dimension_worst float64 dtype: object Missing values: 0
>_ Show codeIn [4]
# Summary statistics for numeric features df_raw.describe().round(3)
| radius_mean | texture_mean | perimeter_mean | area_mean | smoothness_mean | compactness_mean | concavity_mean | concave points_mean | symmetry_mean | fractal_dimension_mean | ... | radius_worst | texture_worst | perimeter_worst | area_worst | smoothness_worst | compactness_worst | concavity_worst | concave points_worst | symmetry_worst | fractal_dimension_worst | |
|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|
| count | 569.000 | 569.000 | 569.000 | 569.000 | 569.000 | 569.000 | 569.000 | 569.000 | 569.000 | 569.000 | ... | 569.000 | 569.000 | 569.000 | 569.000 | 569.000 | 569.000 | 569.000 | 569.000 | 569.000 | 569.000 |
| mean | 14.127 | 19.290 | 91.969 | 654.889 | 0.096 | 0.104 | 0.089 | 0.049 | 0.181 | 0.063 | ... | 16.269 | 25.677 | 107.261 | 880.583 | 0.132 | 0.254 | 0.272 | 0.115 | 0.290 | 0.084 |
| std | 3.524 | 4.301 | 24.299 | 351.914 | 0.014 | 0.053 | 0.080 | 0.039 | 0.027 | 0.007 | ... | 4.833 | 6.146 | 33.603 | 569.357 | 0.023 | 0.157 | 0.209 | 0.066 | 0.062 | 0.018 |
| min | 6.981 | 9.710 | 43.790 | 143.500 | 0.053 | 0.019 | 0.000 | 0.000 | 0.106 | 0.050 | ... | 7.930 | 12.020 | 50.410 | 185.200 | 0.071 | 0.027 | 0.000 | 0.000 | 0.156 | 0.055 |
| 25% | 11.700 | 16.170 | 75.170 | 420.300 | 0.086 | 0.065 | 0.030 | 0.020 | 0.162 | 0.058 | ... | 13.010 | 21.080 | 84.110 | 515.300 | 0.117 | 0.147 | 0.114 | 0.065 | 0.250 | 0.071 |
| 50% | 13.370 | 18.840 | 86.240 | 551.100 | 0.096 | 0.093 | 0.062 | 0.034 | 0.179 | 0.062 | ... | 14.970 | 25.410 | 97.660 | 686.500 | 0.131 | 0.212 | 0.227 | 0.100 | 0.282 | 0.080 |
| 75% | 15.780 | 21.800 | 104.100 | 782.700 | 0.105 | 0.130 | 0.131 | 0.074 | 0.196 | 0.066 | ... | 18.790 | 29.720 | 125.400 | 1084.000 | 0.146 | 0.339 | 0.383 | 0.161 | 0.318 | 0.092 |
| max | 28.110 | 39.280 | 188.500 | 2501.000 | 0.163 | 0.345 | 0.427 | 0.201 | 0.304 | 0.097 | ... | 36.040 | 49.540 | 251.200 | 4254.000 | 0.223 | 1.058 | 1.252 | 0.291 | 0.664 | 0.208 |
>_ Show codeIn [5]
# Class balance
class_counts = df_raw['diagnosis'].value_counts()
class_pct = df_raw['diagnosis'].value_counts(normalize=True) * 100
print('Class distribution:')
for label in ['B', 'M']:
print(f' {label}: {class_counts[label]} ({class_pct[label]:.1f}%)')
print(f'\nClass ratio (M/B): {class_counts["M"] / class_counts["B"]:.3f}')
print('Dataset is moderately imbalanced (63% benign, 37% malignant).')
print('This makes Recall, F1, and ROC-AUC more informative than raw Accuracy.')Class distribution: B: 357 (62.7%) M: 212 (37.3%) Class ratio (M/B): 0.594 Dataset is moderately imbalanced (63% benign, 37% malignant). This makes Recall, F1, and ROC-AUC more informative than raw Accuracy.
Exploratory Data Analysis (EDA)
We visualize the class distribution, feature distributions by class, the correlation structure among features, and a pairplot of the most discriminative features. These plots reveal that malignant tumours tend to have larger, less regular nuclei — which aligns with known pathological characteristics of cancer cells.
>_ Show codeIn [6]
# Encode target for plotting
df = encode_target(df_raw, target_col='diagnosis')
X, y = split_features_target(df, target_col='diagnosis')
feature_names = X.columns.tolist()
# Class distribution bar chart
plt.style.use(PLOT_STYLE)
fig, ax = plt.subplots(figsize=(5, 4))
counts = y.value_counts()
ax.bar(['Benign (0)', 'Malignant (1)'], [counts[0], counts[1]], color=['steelblue', 'tomato'])
ax.set_title('Class Distribution')
ax.set_ylabel('Count')
for i, v in enumerate([counts[0], counts[1]]):
ax.text(i, v + 3, str(v), ha='center', fontweight='bold')
plt.tight_layout()
fig.savefig(os.path.join(PLOTS_DIR, 'class_distribution.png'), dpi=150, bbox_inches='tight')
plt.show()
print('Class distribution: 357 Benign (62.7%), 212 Malignant (37.3%)')
Class distribution: 357 Benign (62.7%), 212 Malignant (37.3%)
>_ Show codeIn [7]
# Feature distributions by class (first 6 mean features)
mean_features = [c for c in feature_names if c.endswith('_mean')][:6]
fig, axes = plt.subplots(2, 3, figsize=(14, 8))
axes = axes.flatten()
for i, feat in enumerate(mean_features):
for label, color, name in [(0, 'steelblue', 'Benign'), (1, 'tomato', 'Malignant')]:
axes[i].hist(df[df['diagnosis'] == label][feat], bins=30, alpha=0.6,
color=color, label=name, density=True)
axes[i].set_title(feat)
axes[i].legend(fontsize=8)
plt.suptitle('Feature Distributions by Class (Mean Features)', fontsize=13, y=1.02)
plt.tight_layout()
fig.savefig(os.path.join(PLOTS_DIR, 'feature_distributions.png'), dpi=150, bbox_inches='tight')
plt.show()
>_ Show codeIn [8]
# Correlation heatmap
fig, ax = plt.subplots(figsize=(16, 13))
corr = df[feature_names].corr()
mask = np.triu(np.ones_like(corr, dtype=bool))
sns.heatmap(corr, mask=mask, cmap='coolwarm', center=0, annot=False,
linewidths=0.5, ax=ax, vmin=-1, vmax=1)
ax.set_title('Feature Correlation Heatmap', fontsize=14)
plt.tight_layout()
fig.savefig(os.path.join(PLOTS_DIR, 'correlation_heatmap.png'), dpi=150, bbox_inches='tight')
plt.show()
print('Many highly correlated feature groups (e.g., radius/perimeter/area).')
print('This multi-collinearity will penalize linear models but is handled well by tree ensembles.')
Many highly correlated feature groups (e.g., radius/perimeter/area). This multi-collinearity will penalize linear models but is handled well by tree ensembles.
>_ Show codeIn [9]
# Pairplot of top 4 discriminative features
top4 = ['radius_mean', 'texture_mean', 'concave points_mean', 'area_mean']
pair_df = df[top4 + ['diagnosis']].copy()
pair_df['diagnosis'] = pair_df['diagnosis'].map({0: 'Benign', 1: 'Malignant'})
g = sns.pairplot(pair_df, hue='diagnosis',
palette={'Benign': 'steelblue', 'Malignant': 'tomato'},
plot_kws={'alpha': 0.5}, diag_kind='hist')
g.fig.suptitle('Pairplot of Top 4 Features', y=1.02, fontsize=13)
g.fig.savefig(os.path.join(PLOTS_DIR, 'pairplot_top_features.png'), dpi=150, bbox_inches='tight')
plt.show()
print('Strong linear separation visible in radius_mean vs concave_points_mean.')
print('This suggests linear models (Logistic Regression) may perform well.')
Strong linear separation visible in radius_mean vs concave_points_mean. This suggests linear models (Logistic Regression) may perform well.
Data Preparation
Train/Test Split
We use an 80/20 stratified split to preserve the class ratio in both sets. With 569 samples, this gives 455 training and 114 test instances.
Preventing Data Leakage
The StandardScaler is fitted exclusively on the training set and then applied to both train and test sets. Fitting on the full dataset — or the test set — would constitute data leakage: the scaler would encode information about the test distribution (mean, variance) into the feature transformation, causing the model to have implicitly "seen" test data during training. This leads to overly optimistic evaluation metrics that do not generalise to truly unseen data. By fitting only on training data, we simulate a realistic deployment scenario where future data is genuinely unknown at training time.
>_ Show codeIn [10]
X_train, X_test, y_train, y_test = get_train_test_split(
X, y, test_size=0.2, random_state=RANDOM_STATE
)
# Fit scaler ONLY on train data — prevents data leakage
scaler = fit_scaler(X_train)
X_train_sc = apply_scaler(scaler, X_train)
X_test_sc = apply_scaler(scaler, X_test)
print(f'Training set: {X_train_sc.shape} | Test set: {X_test_sc.shape}')
print(f'Train class balance: {y_train.mean():.3f} malignant')
print(f'Test class balance: {y_test.mean():.3f} malignant')
print('Stratification successful: class ratio preserved in both splits.')Training set: (455, 30) | Test set: (114, 30) Train class balance: 0.374 malignant Test class balance: 0.368 malignant Stratification successful: class ratio preserved in both splits.
Modeling
Model Selection Rationale
We train four classifiers representing different algorithmic families:
| Model | Rationale |
|---|---|
| Logistic Regression | Interpretable linear baseline; well-suited when features are roughly linearly separable (as the pairplot suggests). Regularisation handles correlated features. |
| Decision Tree | Non-linear, single-tree model. Highly interpretable but prone to overfitting without pruning. Serves as a weak baseline for ensemble comparison. |
| Random Forest | Ensemble of decorrelated trees. Robust to outliers and multicollinearity; captures non-linear interactions. A strong general-purpose baseline. |
| Gradient Boosting | Sequential ensemble that corrects errors of previous trees. Often achieves top performance on tabular data. |
Hyperparameter Tuning
The two ensemble models (Random Forest and Gradient Boosting) are tuned with GridSearchCV using 5-fold cross-validation and ROC-AUC as the scoring metric, which is appropriate for an imbalanced binary classification problem in a clinical context.
>_ Show codeIn [11]
# Define baseline models
models = {
'Logistic Regression': LogisticRegression(max_iter=1000, random_state=RANDOM_STATE),
'Decision Tree': DecisionTreeClassifier(random_state=RANDOM_STATE),
'Random Forest': RandomForestClassifier(n_estimators=100, random_state=RANDOM_STATE),
'Gradient Boosting': GradientBoostingClassifier(n_estimators=100, random_state=RANDOM_STATE),
}
baseline_results = {}
baseline_preds = {}
baseline_probs = {}
for name, model in models.items():
model.fit(X_train_sc, y_train)
y_pred = model.predict(X_test_sc)
y_prob = model.predict_proba(X_test_sc)[:, 1]
metrics = compute_metrics(y_test, y_pred, y_prob)
baseline_results[name] = metrics
baseline_preds[name] = y_pred
baseline_probs[name] = y_prob
print(f'{name}: Acc={metrics["Accuracy"]:.4f}, F1={metrics["F1"]:.4f}, AUC={metrics["ROC-AUC"]:.4f}')Logistic Regression: Acc=0.9649, F1=0.9512, AUC=0.9960 Decision Tree: Acc=0.9298, F1=0.9048, AUC=0.9246
Random Forest: Acc=0.9737, F1=0.9630, AUC=0.9929
Gradient Boosting: Acc=0.9649, F1=0.9500, AUC=0.9947
>_ Show codeIn [12]
# Baseline metrics table
baseline_df = build_metrics_table(baseline_results)
print('Baseline Model Performance:')
baseline_dfBaseline Model Performance:
| Accuracy | Precision | Recall | F1 | ROC-AUC | |
|---|---|---|---|---|---|
| Logistic Regression | 0.9649 | 0.9750 | 0.9286 | 0.9512 | 0.9960 |
| Decision Tree | 0.9298 | 0.9048 | 0.9048 | 0.9048 | 0.9246 |
| Random Forest | 0.9737 | 1.0000 | 0.9286 | 0.9630 | 0.9929 |
| Gradient Boosting | 0.9649 | 1.0000 | 0.9048 | 0.9500 | 0.9947 |
>_ Show codeIn [13]
# GridSearchCV — Random Forest
rf_param_grid = {
'n_estimators': [100, 200],
'max_depth': [None, 5, 10],
'min_samples_split': [2, 5],
}
rf_grid = GridSearchCV(
RandomForestClassifier(random_state=RANDOM_STATE),
rf_param_grid, cv=5, scoring='roc_auc', n_jobs=-1, verbose=0
)
rf_grid.fit(X_train_sc, y_train)
print(f'Best RF params: {rf_grid.best_params_}')
print(f'Best RF CV ROC-AUC: {rf_grid.best_score_:.4f}')Best RF params: {'max_depth': None, 'min_samples_split': 5, 'n_estimators': 200}
Best RF CV ROC-AUC: 0.9906>_ Show codeIn [14]
# GridSearchCV — Gradient Boosting
gb_param_grid = {
'n_estimators': [100, 200],
'learning_rate': [0.05, 0.1],
'max_depth': [3, 5],
}
gb_grid = GridSearchCV(
GradientBoostingClassifier(random_state=RANDOM_STATE),
gb_param_grid, cv=5, scoring='roc_auc', n_jobs=-1, verbose=0
)
gb_grid.fit(X_train_sc, y_train)
print(f'Best GB params: {gb_grid.best_params_}')
print(f'Best GB CV ROC-AUC: {gb_grid.best_score_:.4f}')Best GB params: {'learning_rate': 0.1, 'max_depth': 3, 'n_estimators': 200}
Best GB CV ROC-AUC: 0.9893>_ Show codeIn [15]
# Evaluate tuned models
tuned_models = {
'Random Forest (Tuned)': rf_grid.best_estimator_,
'Gradient Boosting (Tuned)': gb_grid.best_estimator_,
}
tuned_results = {}
tuned_preds = {}
tuned_probs = {}
for name, model in tuned_models.items():
y_pred = model.predict(X_test_sc)
y_prob = model.predict_proba(X_test_sc)[:, 1]
metrics = compute_metrics(y_test, y_pred, y_prob)
tuned_results[name] = metrics
tuned_preds[name] = y_pred
tuned_probs[name] = y_prob
print(f'{name}: Acc={metrics["Accuracy"]:.4f}, F1={metrics["F1"]:.4f}, AUC={metrics["ROC-AUC"]:.4f}')Random Forest (Tuned): Acc=0.9737, F1=0.9630, AUC=0.9950 Gradient Boosting (Tuned): Acc=0.9649, F1=0.9500, AUC=0.9954
>_ Show codeIn [16]
# Combined results
all_results = {**baseline_results, **tuned_results}
all_preds = {**baseline_preds, **tuned_preds}
all_probs = {**baseline_probs, **tuned_probs}
all_df = build_metrics_table(all_results)
print('All Models — Complete Metrics Table:')
all_dfAll Models — Complete Metrics Table:
| Accuracy | Precision | Recall | F1 | ROC-AUC | |
|---|---|---|---|---|---|
| Logistic Regression | 0.9649 | 0.9750 | 0.9286 | 0.9512 | 0.9960 |
| Decision Tree | 0.9298 | 0.9048 | 0.9048 | 0.9048 | 0.9246 |
| Random Forest | 0.9737 | 1.0000 | 0.9286 | 0.9630 | 0.9929 |
| Gradient Boosting | 0.9649 | 1.0000 | 0.9048 | 0.9500 | 0.9947 |
| Random Forest (Tuned) | 0.9737 | 1.0000 | 0.9286 | 0.9630 | 0.9950 |
| Gradient Boosting (Tuned) | 0.9649 | 1.0000 | 0.9048 | 0.9500 | 0.9954 |
>_ Show codeIn [17]
# Select best model by ROC-AUC
best_model_name = all_df['ROC-AUC'].idxmax()
print(f'Best model by ROC-AUC: {best_model_name}')
print(f'ROC-AUC: {all_df.loc[best_model_name, "ROC-AUC"]:.4f}')
model_obj_map = {
'Logistic Regression': models['Logistic Regression'],
'Decision Tree': models['Decision Tree'],
'Random Forest': models['Random Forest'],
'Gradient Boosting': models['Gradient Boosting'],
'Random Forest (Tuned)': rf_grid.best_estimator_,
'Gradient Boosting (Tuned)': gb_grid.best_estimator_,
}
best_model = model_obj_map[best_model_name]
best_probs = all_probs[best_model_name]
best_preds = all_preds[best_model_name]
# Save best model
with open(os.path.join(MODELS_DIR, 'best_model.pkl'), 'wb') as f:
pickle.dump(best_model, f)
print(f'Best model saved to outputs/models/best_model.pkl')Best model by ROC-AUC: Logistic Regression ROC-AUC: 0.9960 Best model saved to outputs/models/best_model.pkl
Evaluation
We evaluate all six models (4 baseline + 2 tuned) using five metrics:
- Accuracy: overall correctness
- Precision: proportion of positive predictions that are truly positive (minimises false positives)
- Recall: proportion of true positives correctly identified (minimises false negatives — critical in cancer screening)
- F1: harmonic mean of precision and recall
- ROC-AUC: discrimination ability across all thresholds; robust to class imbalance
For clinical applications, Recall and ROC-AUC are the primary metrics of interest.
>_ Show codeIn [18]
# Display complete metrics table
print('Complete Metrics — All Models (Baseline + Tuned):')
all_df.style.highlight_max(color='lightgreen', axis=0)Complete Metrics — All Models (Baseline + Tuned):
| Accuracy | Precision | Recall | F1 | ROC-AUC | |
|---|---|---|---|---|---|
| Logistic Regression | 0.964900 | 0.975000 | 0.928600 | 0.951200 | 0.996000 |
| Decision Tree | 0.929800 | 0.904800 | 0.904800 | 0.904800 | 0.924600 |
| Random Forest | 0.973700 | 1.000000 | 0.928600 | 0.963000 | 0.992900 |
| Gradient Boosting | 0.964900 | 1.000000 | 0.904800 | 0.950000 | 0.994700 |
| Random Forest (Tuned) | 0.973700 | 1.000000 | 0.928600 | 0.963000 | 0.995000 |
| Gradient Boosting (Tuned) | 0.964900 | 1.000000 | 0.904800 | 0.950000 | 0.995400 |
>_ Show codeIn [19]
# Confusion matrices for all models
fig, axes = plt.subplots(2, 3, figsize=(16, 10))
axes = axes.flatten()
model_names = list(all_results.keys())
class_names = ['Benign (0)', 'Malignant (1)']
for idx, name in enumerate(model_names):
cm = confusion_matrix(y_test, all_preds[name])
sns.heatmap(cm, annot=True, fmt='d', cmap='Blues',
xticklabels=class_names, yticklabels=class_names, ax=axes[idx])
axes[idx].set_title(name, fontsize=10)
axes[idx].set_xlabel('Predicted')
axes[idx].set_ylabel('True')
plt.suptitle('Confusion Matrices — All Models', fontsize=14, y=1.02)
plt.tight_layout()
fig.savefig(os.path.join(PLOTS_DIR, 'confusion_matrices_all.png'), dpi=150, bbox_inches='tight')
plt.show()
# Also save individual confusion matrices
for name in model_names:
safe = name.lower().replace(' ', '_').replace('(', '').replace(')', '')
plot_confusion_matrix(y_test, all_preds[name], model_name=name,
save_path=os.path.join(PLOTS_DIR, f'confusion_matrix_{safe}.png'))
plt.close('all')
print('Confusion matrices saved.')
Confusion matrices saved.
>_ Show codeIn [20]
# ROC curves — all models
plt.style.use(PLOT_STYLE)
fig, ax = plt.subplots(figsize=(8, 6))
colors = ['steelblue', 'darkorange', 'green', 'red', 'purple', 'brown']
for (name, y_prob), color in zip(all_probs.items(), colors):
fpr, tpr, _ = roc_curve(y_test, y_prob)
auc = roc_auc_score(y_test, y_prob)
ls = '--' if 'Tuned' in name else '-'
ax.plot(fpr, tpr, lw=2, linestyle=ls, color=color, label=f'{name} (AUC={auc:.4f})')
ax.plot([0, 1], [0, 1], 'k--', lw=1, label='Random classifier')
ax.set_xlabel('False Positive Rate')
ax.set_ylabel('True Positive Rate')
ax.set_title('ROC Curves — All Models')
ax.legend(loc='lower right', fontsize=8)
plt.tight_layout()
fig.savefig(os.path.join(PLOTS_DIR, 'roc_curves.png'), dpi=150, bbox_inches='tight')
plt.show()
Threshold Analysis
Clinical Context for Threshold Selection
The default 0.5 classification threshold optimises for overall accuracy, but in clinical practice the threshold should reflect the cost asymmetry between error types:
Screening context (community-level testing): A false negative (missed cancer) has severe consequences — delayed treatment. We lower the threshold to 0.3, accepting more false positives (benign recalls for follow-up) to catch virtually all malignancies. The motto: "When in doubt, flag it."
Confirmatory context (specialist review before invasive procedure): A false positive causes unnecessary biopsy with associated risk and cost. A higher threshold (e.g., 0.6–0.7) may be appropriate to increase specificity.
We analyse how precision and recall trade off as the threshold varies, and demonstrate the effect of a clinically motivated threshold of 0.3 on the confusion matrix.
>_ Show codeIn [21]
# Precision-Recall curve
plt.style.use(PLOT_STYLE)
precision_vals, recall_vals, thresholds = precision_recall_curve(y_test, best_probs)
fig, ax = plt.subplots(figsize=(7, 5))
ax.plot(recall_vals, precision_vals, lw=2, color='steelblue')
ax.set_xlabel('Recall')
ax.set_ylabel('Precision')
ax.set_title(f'Precision-Recall Curve — {best_model_name}')
ax.set_xlim([0, 1])
ax.set_ylim([0, 1.05])
plt.tight_layout()
fig.savefig(os.path.join(PLOTS_DIR, 'precision_recall_curve.png'), dpi=150, bbox_inches='tight')
plt.show()
>_ Show codeIn [22]
# Precision and Recall vs Threshold
plt.style.use(PLOT_STYLE)
thresholds_full = np.append(thresholds, 1.0)
fig, ax = plt.subplots(figsize=(9, 5))
ax.plot(thresholds_full, precision_vals, lw=2, label='Precision', color='darkorange')
ax.plot(thresholds_full, recall_vals, lw=2, label='Recall', color='steelblue')
ax.axvline(x=0.3, color='red', linestyle='--', lw=1.5, label='Screening threshold (0.3)')
ax.axvline(x=0.5, color='green', linestyle='--', lw=1.5, label='Default threshold (0.5)')
ax.set_xlabel('Classification Threshold')
ax.set_ylabel('Score')
ax.set_title(f'Precision & Recall vs Threshold — {best_model_name}')
ax.legend()
ax.set_xlim([0, 1])
ax.set_ylim([0, 1.05])
plt.tight_layout()
fig.savefig(os.path.join(PLOTS_DIR, 'threshold_analysis.png'), dpi=150, bbox_inches='tight')
plt.show()
>_ Show codeIn [23]
# Clinical threshold recommendation
clinical_threshold = 0.3
best_preds_clinical = (best_probs >= clinical_threshold).astype(int)
metrics_default = compute_metrics(y_test, best_preds, best_probs)
metrics_clinical = compute_metrics(y_test, best_preds_clinical, best_probs)
print(f'=== {best_model_name} ===')
print(f'\nDefault threshold (0.5):')
for k, v in metrics_default.items():
print(f' {k}: {v:.4f}')
print(f'\nScreening threshold (0.3):')
for k, v in metrics_clinical.items():
print(f' {k}: {v:.4f}')
print('\nAt threshold 0.3: Recall improves at the cost of slightly lower Precision.')
print('Recommended for population screening to minimise missed malignancies.')=== Logistic Regression === Default threshold (0.5): Accuracy: 0.9649 Precision: 0.9750 Recall: 0.9286 F1: 0.9512 ROC-AUC: 0.9960 Screening threshold (0.3): Accuracy: 0.9825 Precision: 0.9762 Recall: 0.9762 F1: 0.9762 ROC-AUC: 0.9960 At threshold 0.3: Recall improves at the cost of slightly lower Precision. Recommended for population screening to minimise missed malignancies.
>_ Show codeIn [24]
# Confusion matrix at clinical threshold
fig, axes = plt.subplots(1, 2, figsize=(11, 4))
class_names = ['Benign (0)', 'Malignant (1)']
for ax, preds, title in [
(axes[0], best_preds, f'{best_model_name}\n(Default threshold 0.5)'),
(axes[1], best_preds_clinical, f'{best_model_name}\n(Screening threshold 0.3)'),
]:
cm = confusion_matrix(y_test, preds)
sns.heatmap(cm, annot=True, fmt='d', cmap='Blues',
xticklabels=class_names, yticklabels=class_names, ax=ax)
ax.set_title(title, fontsize=10)
ax.set_xlabel('Predicted')
ax.set_ylabel('True')
plt.tight_layout()
fig.savefig(os.path.join(PLOTS_DIR, 'confusion_matrix_clinical_threshold.png'), dpi=150, bbox_inches='tight')
plt.show()
Interpretability — SHAP Analysis
To understand which features drive predictions, we use SHAP (SHapley Additive exPlanations), a game-theoretic framework that assigns each feature a contribution to the model's output for every prediction.
Why SHAP?
SHAP values are model-agnostic in principle but have efficient implementations for tree-based models (TreeExplainer) and linear models (LinearExplainer). They satisfy important axioms (local accuracy, missingness, consistency) that make them a theoretically grounded interpretability method.
Clinical Interpretation
Features with large mean |SHAP| values are the most influential predictors of malignancy. We expect concave_points_worst, radius_worst, and perimeter_worst to dominate — larger, more irregular nuclei are hallmarks of malignant cells in cytology.
>_ Show codeIn [25]
# SHAP analysis for best model
print(f'Computing SHAP values for: {best_model_name}')
if 'Gradient Boosting' in best_model_name or 'Random Forest' in best_model_name:
explainer = shap.TreeExplainer(best_model)
shap_values = explainer.shap_values(X_test_sc)
if isinstance(shap_values, list):
sv = shap_values[1]
else:
sv = shap_values
else:
explainer = shap.LinearExplainer(best_model, X_train_sc, feature_perturbation='interventional')
shap_values = explainer.shap_values(X_test_sc)
sv = shap_values
print(f'SHAP values shape: {sv.shape}')Computing SHAP values for: Logistic Regression SHAP values shape: (114, 30)
>_ Show codeIn [26]
# SHAP summary beeswarm plot
plt.figure(figsize=(10, 8))
shap.summary_plot(sv, X_test_sc, feature_names=feature_names, show=False)
plt.title(f'SHAP Summary Plot (Beeswarm) — {best_model_name}', fontsize=13)
plt.tight_layout()
plt.savefig(os.path.join(PLOTS_DIR, 'shap_summary.png'), dpi=150, bbox_inches='tight')
plt.show()
print('SHAP beeswarm: each dot is one test sample. Red = high feature value, blue = low.')
print('Positive SHAP value -> pushes prediction toward Malignant.')
SHAP beeswarm: each dot is one test sample. Red = high feature value, blue = low. Positive SHAP value -> pushes prediction toward Malignant.
>_ Show codeIn [27]
# SHAP bar plot — mean absolute SHAP values (global feature importance)
plt.figure(figsize=(9, 7))
shap.summary_plot(sv, X_test_sc, feature_names=feature_names, plot_type='bar', show=False)
plt.title(f'SHAP Global Feature Importance (Mean |SHAP|) — {best_model_name}', fontsize=12)
plt.tight_layout()
plt.savefig(os.path.join(PLOTS_DIR, 'shap_bar.png'), dpi=150, bbox_inches='tight')
plt.show()
>_ Show codeIn [28]
# Top 10 features by mean |SHAP|
mean_abs_shap = np.abs(sv).mean(axis=0)
shap_importance = pd.Series(mean_abs_shap, index=feature_names).sort_values(ascending=False)
print('Top 10 features by mean |SHAP| value:')
print(shap_importance.head(10).round(4).to_string())
print('\nClinical interpretation:')
print(' - concave_points_worst / concavity_worst: measure nuclear irregularity — higher in malignant cells')
print(' - radius_worst / perimeter_worst / area_worst: nuclear size — malignant nuclei are larger')
print(' - texture_worst: nuclear texture (standard deviation of grayscale values) — coarser in cancer')
print('These findings are consistent with established cytological criteria for malignancy.')Top 10 features by mean |SHAP| value: texture_worst 1.2507 radius_se 0.7611 symmetry_worst 0.7406 concave points_mean 0.7228 radius_worst 0.7125 concavity_worst 0.6937 compactness_se 0.6887 area_worst 0.6846 concave points_worst 0.5749 perimeter_worst 0.5614 Clinical interpretation: - concave_points_worst / concavity_worst: measure nuclear irregularity — higher in malignant cells - radius_worst / perimeter_worst / area_worst: nuclear size — malignant nuclei are larger - texture_worst: nuclear texture (standard deviation of grayscale values) — coarser in cancer These findings are consistent with established cytological criteria for malignancy.
Conclusions
Key Findings
All four classifiers achieved strong performance on this dataset, reflecting the high signal-to-noise ratio of the WDBC features. The complete results are summarised below:
| Model | Accuracy | Precision | Recall | F1 | ROC-AUC |
|---|---|---|---|---|---|
| Logistic Regression | 0.9649 | 0.9750 | 0.9286 | 0.9512 | 0.9960 |
| Decision Tree | 0.9298 | 0.9048 | 0.9048 | 0.9048 | 0.9246 |
| Random Forest | 0.9737 | 1.0000 | 0.9286 | 0.9630 | 0.9929 |
| Gradient Boosting | 0.9649 | 1.0000 | 0.9048 | 0.9500 | 0.9947 |
| Random Forest (Tuned) | 0.9737 | 1.0000 | 0.9286 | 0.9630 | 0.9950 |
| Gradient Boosting (Tuned) | 0.9649 | 1.0000 | 0.9048 | 0.9500 | 0.9954 |
Why Logistic Regression Wins
Logistic Regression achieved the highest ROC-AUC (0.9960), surpassing even tuned ensemble methods. This outcome is explained by the dataset's characteristics:
Near-linear separability: The pairplot and EDA showed that malignant and benign samples are largely linearly separable in the feature space — particularly in
concave_points_meanvsradius_mean. Logistic Regression is optimally designed for this structure.L2 regularisation handles multicollinearity: Many features are highly correlated (e.g., radius, perimeter, area). Logistic Regression's L2 penalty effectively distributes weight across correlated features, resulting in stable and well-calibrated probability estimates — which directly improves AUC.
Sample size: With only 455 training samples, complex models like Gradient Boosting can overfit the training distribution. Logistic Regression's simplicity is a virtue at this scale.
Well-calibrated probabilities: Tree ensembles produce probabilities that are often poorly calibrated (scores cluster near 0 and 1). Logistic Regression produces better-calibrated scores, leading to a smoother and higher ROC curve.
Clinical Threshold Recommendation
For a screening application, we recommend a threshold of 0.3, which maximises Recall (catching more malignancies at the cost of some additional false positives). In a confirmatory / pre-surgical context, the default threshold of 0.5 or higher is appropriate to minimise unnecessary procedures.
Most Predictive Features (SHAP)
The SHAP analysis confirms that concave_points_worst, radius_worst, and perimeter_worst are the strongest predictors of malignancy. These correspond directly to well-established cytological criteria: malignant cells have larger, more irregular, and more concave nuclei.
Limitations
- Dataset size: 569 samples is relatively small; larger multi-site datasets are needed before clinical deployment.
- Single modality: The dataset is based solely on FNA measurements; integration with imaging (mammography, ultrasound) and clinical features would improve diagnostic value.
- No prospective validation: Results are based on a retrospective dataset; prospective validation is required.
- Class imbalance: While moderate (63/37), a real-world screening population would have a much lower malignancy prevalence, which would affect precision substantially.
Next Steps
- Calibration: Apply Platt scaling or isotonic regression to improve probability calibration for the ensemble models.
- Cross-dataset validation: Validate on independent FNA datasets from different institutions.
- Feature engineering: Explore ratios and polynomial interactions between nuclear measurements.
- Cost-sensitive learning: Incorporate explicit misclassification costs (e.g., 5x cost for false negatives) directly into model training.
- Fairness analysis: Examine whether model performance is consistent across demographic subgroups.
>_ Show codeIn [29]
# Final summary printout
print('=' * 60)
print('FINAL RESULTS SUMMARY')
print('=' * 60)
print(f'Best model: {best_model_name}')
print(f'Test set metrics:')
for k, v in all_results[best_model_name].items():
print(f' {k}: {v:.4f}')
print(f'\nAll outputs saved to outputs/plots/ and outputs/models/')
print(f'Best model pickle: outputs/models/best_model.pkl')
print('=' * 60)============================================================ FINAL RESULTS SUMMARY ============================================================ Best model: Logistic Regression Test set metrics: Accuracy: 0.9649 Precision: 0.9750 Recall: 0.9286 F1: 0.9512 ROC-AUC: 0.9960 All outputs saved to outputs/plots/ and outputs/models/ Best model pickle: outputs/models/best_model.pkl ============================================================
