Handwritten-digit recogniser
scikit-learn ships with about 1,800 tiny scans of handwritten digits, each 8 by 8 pixels. You'll draw a few of them as text so you can see what the model sees, train a k-nearest neighbours model, and look at the digits it gets wrong.
Skills you'll practise
- Images as numbers
- k-nearest neighbours
- Accuracy
- Reading mistakes
Steps
Step 1: Look at a digit
Load the digits and print one 8×8 image as text, using a character for each shade of grey.
Show a hintHide the hint
Pixel values run from 0 (blank) to 16 (darkest). Picking a character by
pixel // 5gives four shades.Step 2: Split the scans
Use
digits.data, where each image is flattened into a row of 64 numbers, and hold back a quarter for testing.Step 3: Train kNN
Fit a
KNeighborsClassifierwith three neighbours on the training scans.Show a hintHide the hint
kNN labels a new digit by finding the training digits whose 64 pixel values are closest to it.
Step 4: Measure it
Print the accuracy on the test scans.
Step 5: Look at the mistakes
Find the test digits the model got wrong and draw one, with what it guessed and what it really was.
Show a hintHide the hint
wrong = np.where(predictions != y_test)[0]gives the positions of the mistakes.
Starter code
It already runs. The TODO comments mark where to start.
from sklearn.datasets import load_digits
digits = load_digits()
SHADES = " .:#" # blank, light, medium, dark
def draw(image):
for row in image:
print("".join(SHADES[int(pixel) // 5] for pixel in row))
print("This is a", digits.target[0])
draw(digits.images[0])
# TODO: split digits.data and digits.target, train a
# KNeighborsClassifier(n_neighbors=3), and print its accuracy.
Example solution
One way to finish it. Have a go first; yours doesn't need to match.
Reveal the solutionHide the solution
import numpy as np
from sklearn.datasets import load_digits
from sklearn.model_selection import train_test_split
from sklearn.neighbors import KNeighborsClassifier
from sklearn.metrics import accuracy_score
digits = load_digits()
SHADES = " .:#" # blank, light, medium, dark
def draw(image):
for row in image:
print("".join(SHADES[int(pixel) // 5] for pixel in row))
X_train, X_test, y_train, y_test = train_test_split(
digits.data, digits.target, test_size=0.25, random_state=0, stratify=digits.target
)
model = KNeighborsClassifier(n_neighbors=3)
model.fit(X_train, y_train)
predictions = model.predict(X_test)
print(f"Accuracy on unseen digits: {accuracy_score(y_test, predictions):.1%}")
wrong = np.where(predictions != y_test)[0]
print(f"It got {len(wrong)} of {len(y_test)} wrong.")
if len(wrong) > 0:
first = wrong[0]
print()
print(f"It guessed {predictions[first]}, but this is a {y_test[first]}:")
draw(X_test[first].reshape(8, 8))
Stretch goal
Use matplotlib to show a grid of the digits it got wrong with plt.imshow(image, cmap="gray_r"), titled with the guess and the true label.