Skip to main content

Palmer Penguins kNN Classification (Python)

This Quarto document runs the Python translation of kNN_example_penguins.qmd through reticulate.

library(reticulate)
py_require(c("pandas", "numpy", "scikit-learn", "matplotlib",
             "seaborn", "palmerpenguins"))

Load and inspect the data

from pathlib import Path

import matplotlib.pyplot as plt
import pandas as pd
import seaborn as sns
from sklearn.metrics import accuracy_score, ConfusionMatrixDisplay, confusion_matrix
from sklearn.model_selection import train_test_split
from sklearn.neighbors import KNeighborsClassifier

DATA_FILE = Path("penguins.csv")

if DATA_FILE.exists():
    penguins = pd.read_csv(DATA_FILE)
else:
    try:
        from palmerpenguins import load_penguins
        penguins = load_penguins()
    except ImportError:
        penguins = sns.load_dataset("penguins")

penguins = penguins.drop(columns=["island", "sex"])
print(penguins.head())
  species  bill_length_mm  bill_depth_mm  flipper_length_mm  body_mass_g
0  Adelie            39.1           18.7              181.0       3750.0
1  Adelie            39.5           17.4              186.0       3800.0
2  Adelie            40.3           18.0              195.0       3250.0
3  Adelie             NaN            NaN                NaN          NaN
4  Adelie            36.7           19.3              193.0       3450.0
print(penguins["species"].value_counts().sort_index())
species
Adelie       152
Chinstrap     68
Gentoo       124
Name: count, dtype: int64
print(penguins.isna().sum())
species              0
bill_length_mm       2
bill_depth_mm        2
flipper_length_mm    2
body_mass_g          2
dtype: int64

Visualize the data

sns.scatterplot(data=penguins, x="bill_length_mm", y="bill_depth_mm")
plt.tight_layout()
plt.show()

sns.scatterplot(
    data=penguins, x="bill_length_mm", y="bill_depth_mm", hue="species"
)
plt.tight_layout()
plt.show()

g = sns.FacetGrid(penguins, col="species")
g.map_dataframe(sns.scatterplot, x="bill_length_mm", y="bill_depth_mm")

g.tight_layout()

plt.show()

Fit and evaluate the kNN model

The original tutorial uses an unweighted kNN model with k = 4. Python’s scikit-learn requires complete predictor rows, so missing rows are removed before fitting.

penguins_complete = penguins.dropna().copy()

X = penguins_complete.drop(columns="species")
y = penguins_complete["species"]
X_train, X_test, y_train, y_test = train_test_split(
    X, y, test_size=0.20, stratify=y, random_state=123
)

knn_model = KNeighborsClassifier(n_neighbors=4, weights="uniform")
knn_model.fit(X_train, y_train)
KNeighborsClassifier(n_neighbors=4)
In a Jupyter environment, please rerun this cell to show the HTML representation or trust the notebook.
On GitHub, the HTML representation is unable to render, please try loading this page with nbviewer.org.
def evaluate(X_data, y_data, label):
    predictions = knn_model.predict(X_data)
    cm = confusion_matrix(y_data, predictions, labels=knn_model.classes_)
    print(f"{label} accuracy: {accuracy_score(y_data, predictions):.4f}")
    print(pd.DataFrame(cm, index=knn_model.classes_, columns=knn_model.classes_))
    ConfusionMatrixDisplay(cm, display_labels=knn_model.classes_).plot(cmap="Blues")
    plt.title(f"{label} confusion matrix")
    plt.tight_layout()
    plt.show()
    return predictions, cm

train_predictions, train_cm = evaluate(X_train, y_train, "Training")
Training accuracy: 0.8681
           Adelie  Chinstrap  Gentoo
Adelie        118          1       2
Chinstrap      27         25       2
Gentoo          4          0      94

test_predictions, test_cm = evaluate(X_test, y_test, "Testing")
Testing accuracy: 0.8841
           Adelie  Chinstrap  Gentoo
Adelie         30          0       0
Chinstrap       6          8       0
Gentoo          1          1      23