15: Classification
- Page ID
- 50860
\( \newcommand{\vecs}[1]{\overset { \scriptstyle \rightharpoonup} {\mathbf{#1}} } \)
\( \newcommand{\vecd}[1]{\overset{-\!-\!\rightharpoonup}{\vphantom{a}\smash {#1}}} \)
\( \newcommand{\dsum}{\displaystyle\sum\limits} \)
\( \newcommand{\dint}{\displaystyle\int\limits} \)
\( \newcommand{\dlim}{\displaystyle\lim\limits} \)
\( \newcommand{\id}{\mathrm{id}}\) \( \newcommand{\Span}{\mathrm{span}}\)
( \newcommand{\kernel}{\mathrm{null}\,}\) \( \newcommand{\range}{\mathrm{range}\,}\)
\( \newcommand{\RealPart}{\mathrm{Re}}\) \( \newcommand{\ImaginaryPart}{\mathrm{Im}}\)
\( \newcommand{\Argument}{\mathrm{Arg}}\) \( \newcommand{\norm}[1]{\| #1 \|}\)
\( \newcommand{\inner}[2]{\langle #1, #2 \rangle}\)
\( \newcommand{\Span}{\mathrm{span}}\)
\( \newcommand{\id}{\mathrm{id}}\)
\( \newcommand{\Span}{\mathrm{span}}\)
\( \newcommand{\kernel}{\mathrm{null}\,}\)
\( \newcommand{\range}{\mathrm{range}\,}\)
\( \newcommand{\RealPart}{\mathrm{Re}}\)
\( \newcommand{\ImaginaryPart}{\mathrm{Im}}\)
\( \newcommand{\Argument}{\mathrm{Arg}}\)
\( \newcommand{\norm}[1]{\| #1 \|}\)
\( \newcommand{\inner}[2]{\langle #1, #2 \rangle}\)
\( \newcommand{\Span}{\mathrm{span}}\) \( \newcommand{\AA}{\unicode[.8,0]{x212B}}\)
\( \newcommand{\vectorA}[1]{\vec{#1}} % arrow\)
\( \newcommand{\vectorAt}[1]{\vec{\text{#1}}} % arrow\)
\( \newcommand{\vectorB}[1]{\overset { \scriptstyle \rightharpoonup} {\mathbf{#1}} } \)
\( \newcommand{\vectorC}[1]{\textbf{#1}} \)
\( \newcommand{\vectorD}[1]{\overrightarrow{#1}} \)
\( \newcommand{\vectorDt}[1]{\overrightarrow{\text{#1}}} \)
\( \newcommand{\vectE}[1]{\overset{-\!-\!\rightharpoonup}{\vphantom{a}\smash{\mathbf {#1}}}} \)
\( \newcommand{\vecs}[1]{\overset { \scriptstyle \rightharpoonup} {\mathbf{#1}} } \)
\(\newcommand{\longvect}{\overrightarrow}\)
\( \newcommand{\vecd}[1]{\overset{-\!-\!\rightharpoonup}{\vphantom{a}\smash {#1}}} \)
\(\newcommand{\avec}{\mathbf a}\) \(\newcommand{\bvec}{\mathbf b}\) \(\newcommand{\cvec}{\mathbf c}\) \(\newcommand{\dvec}{\mathbf d}\) \(\newcommand{\dtil}{\widetilde{\mathbf d}}\) \(\newcommand{\evec}{\mathbf e}\) \(\newcommand{\fvec}{\mathbf f}\) \(\newcommand{\nvec}{\mathbf n}\) \(\newcommand{\pvec}{\mathbf p}\) \(\newcommand{\qvec}{\mathbf q}\) \(\newcommand{\svec}{\mathbf s}\) \(\newcommand{\tvec}{\mathbf t}\) \(\newcommand{\uvec}{\mathbf u}\) \(\newcommand{\vvec}{\mathbf v}\) \(\newcommand{\wvec}{\mathbf w}\) \(\newcommand{\xvec}{\mathbf x}\) \(\newcommand{\yvec}{\mathbf y}\) \(\newcommand{\zvec}{\mathbf z}\) \(\newcommand{\rvec}{\mathbf r}\) \(\newcommand{\mvec}{\mathbf m}\) \(\newcommand{\zerovec}{\mathbf 0}\) \(\newcommand{\onevec}{\mathbf 1}\) \(\newcommand{\real}{\mathbb R}\) \(\newcommand{\twovec}[2]{\left[\begin{array}{r}#1 \\ #2 \end{array}\right]}\) \(\newcommand{\ctwovec}[2]{\left[\begin{array}{c}#1 \\ #2 \end{array}\right]}\) \(\newcommand{\threevec}[3]{\left[\begin{array}{r}#1 \\ #2 \\ #3 \end{array}\right]}\) \(\newcommand{\cthreevec}[3]{\left[\begin{array}{c}#1 \\ #2 \\ #3 \end{array}\right]}\) \(\newcommand{\fourvec}[4]{\left[\begin{array}{r}#1 \\ #2 \\ #3 \\ #4 \end{array}\right]}\) \(\newcommand{\cfourvec}[4]{\left[\begin{array}{c}#1 \\ #2 \\ #3 \\ #4 \end{array}\right]}\) \(\newcommand{\fivevec}[5]{\left[\begin{array}{r}#1 \\ #2 \\ #3 \\ #4 \\ #5 \\ \end{array}\right]}\) \(\newcommand{\cfivevec}[5]{\left[\begin{array}{c}#1 \\ #2 \\ #3 \\ #4 \\ #5 \\ \end{array}\right]}\) \(\newcommand{\mattwo}[4]{\left[\begin{array}{rr}#1 \amp #2 \\ #3 \amp #4 \\ \end{array}\right]}\) \(\newcommand{\laspan}[1]{\text{Span}\{#1\}}\) \(\newcommand{\bcal}{\cal B}\) \(\newcommand{\ccal}{\cal C}\) \(\newcommand{\scal}{\cal S}\) \(\newcommand{\wcal}{\cal W}\) \(\newcommand{\ecal}{\cal E}\) \(\newcommand{\coords}[2]{\left\{#1\right\}_{#2}}\) \(\newcommand{\gray}[1]{\color{gray}{#1}}\) \(\newcommand{\lgray}[1]{\color{lightgray}{#1}}\) \(\newcommand{\rank}{\operatorname{rank}}\) \(\newcommand{\row}{\text{Row}}\) \(\newcommand{\col}{\text{Col}}\) \(\renewcommand{\row}{\text{Row}}\) \(\newcommand{\nul}{\text{Nul}}\) \(\newcommand{\var}{\text{Var}}\) \(\newcommand{\corr}{\text{corr}}\) \(\newcommand{\len}[1]{\left|#1\right|}\) \(\newcommand{\bbar}{\overline{\bvec}}\) \(\newcommand{\bhat}{\widehat{\bvec}}\) \(\newcommand{\bperp}{\bvec^\perp}\) \(\newcommand{\xhat}{\widehat{\xvec}}\) \(\newcommand{\vhat}{\widehat{\vvec}}\) \(\newcommand{\uhat}{\widehat{\uvec}}\) \(\newcommand{\what}{\widehat{\wvec}}\) \(\newcommand{\Sighat}{\widehat{\Sigma}}\) \(\newcommand{\lt}{<}\) \(\newcommand{\gt}{>}\) \(\newcommand{\amp}{&}\) \(\definecolor{fillinmathshade}{gray}{0.9}\)In the previous chapter, we learned how regression predicts numerical values, such as miles per gallon or delivery times. In this chapter, we will learn how to predict categories instead of numbers. Classification algorithms determine which category a new data point most likely belongs to, by comparing it with previously labeled data.
In this chapter we cover a simple and intuitive classification method: the nearest neighbor algorithm. We will work with datasets in which each observation belongs to one of only two categories. This is called binary classification.
Gather Data
To investigate what a nearest neighbor is, we work with the cdc dataset that we used in previous chapters. We will determine whether the height and weight of a person can be used to predict the person's gender.
We read in data from the cdc.csv file and inspect it.
| genhlth | exerany | hlthplan | smoke100 | height | weight | wtdesire | age | gender | |
|---|---|---|---|---|---|---|---|---|---|
| 0 | good | 0 | 1 | 0 | 70 | 175 | 175 | 77 | m |
| 1 | good | 0 | 1 | 1 | 64 | 125 | 115 | 33 | f |
| 2 | good | 1 | 1 | 1 | 60 | 105 | 105 | 49 | f |
| 3 | good | 1 | 1 | 0 | 66 | 132 | 124 | 42 | f |
| 4 | very good | 0 | 1 | 0 | 61 | 150 | 130 | 55 | f |
Number of rows: 20000
Prepare Data
We create a new DataFrame from the columns height, weight, and gender.
| height | weight | gender | |
|---|---|---|---|
| 0 | 70 | 175 | m |
| 1 | 64 | 125 | f |
| 2 | 60 | 105 | f |
| 3 | 66 | 132 | f |
| 4 | 61 | 150 | f |
Because the classification algorithm works with numerical data, we need to represent the gender categories as numbers instead of letters. We will replace m with 0 and f with 1.
We arbitrarily assign m to 0 and f to 1, but we can just as easily assign the 0 and 1 values the opposite way. The numbers do not mean that one category is larger or more important than the other. They simply provide two distinct labels that the algorithm can use to identify the two classes.
| height | weight | gender | |
|---|---|---|---|
| 0 | 70 | 175 | 0 |
| 1 | 64 | 125 | 1 |
| 2 | 60 | 105 | 1 |
| 3 | 66 | 132 | 1 |
| 4 | 61 | 150 | 1 |
The gender feature has two categories or two classes: 0 or 1. In a classification problem, the possible output categories are called classes. Therefore, the gender feature has two classes: class 0 and class 1.
We inspect the height and weight data by finding their maximum and minimum.
Maximum height: 93
Minimum height: 48
Maximum weight: 500
Minimum weight: 68
The height and weight are measured on different scales. Height ranges from 48 to 93 inches, while weight ranges from about 68 to 500 pounds. Because the nearest neighbor algorithm calculates the distance between data points, the larger values for weight would have a much greater influence on the distance than the smaller values for height.
Therefore, to give each feature equal importance, we convert both height and weight to standard units. After this conversion, both features are measured on the same scale, so neither one dominates the distance calculation.
We bring in the standard_units function from Chapter 12, then we use the function to convert the data.
| height | weight | gender | |
|---|---|---|---|
| 0 | 0.682792 | 0.132661 | 0 |
| 1 | -0.771453 | -1.114845 | 1 |
| 2 | -1.740950 | -1.613847 | 1 |
| 3 | -0.286704 | -0.940194 | 1 |
| 4 | -1.498576 | -0.491092 | 1 |
Exploratory Data Analysis (EDA)
First we look at the relationship between height and weight by plotting them in a scatterplot.
There is a positive correlation between height and weight. As the height increases, the weight tends to also increase. However, the points are fairly spread out, so people with the same height can have a wide range of weights. As a result, height alone cannot predict a person's weight with high accuracy.
In addition, because the dataset contains 20,000 observations, the scatterplot is so crowded that it is difficult to see any overall pattern. For illustration, we randomly select 500 observations from the dataset. This smaller sample still shows the relationship between the variables while making the individual data points visible.
The nearest neighbor algorithm can later be applied to the full dataset, but we for now use a smaller sample to make the scatterplots and decision boundaries much easier to visualize.
The smaller sample preserves the overall pattern of the data. Although there are fewer points, the scatterplot still shows the same positive relationship between height and weight as the full dataset.
The scatterplot above doesn't tell us which data points belong to class 0 (male) and which belong to class 1 (female). To distinguish the two classes, we use groupby to separate the data into two groups based on the gender. We then plot each group using a different color, making it easier to see the two categories.
Now we can see a pattern with the two classes:
- Males (class 0) generally have larger heights and weights. Most of the male data points are above 0 on the height axis and above −0.5 on the weight axis.
- Females (class 1) generally have smaller heights and weights. Most of the female data points lie between −2 and +1 on the height axis and below 0 on the weight axis.
- The two classes overlap considerably. Some males have similar heights and weights as some females, and many of the blue and orange points appear in the same regions of the scatterplot. Therefore it is not always possible to determine a person's gender from height and weight alone.
Because the two classes overlap, we need a method to decide which class a new observation most likely belongs to. The nearest neighbor algorithm does this by looking at the classes of nearby observations.
Nearest Neighbor
Suppose we have a new person named Kai whose height is −1.2 standard units and whose weight is −0.8 standard units. We would like to predict whether Kai belongs to class 0 (male) or class 1 (female).
We can plot Kai's height and weight in the scatterplot to see where this new data point is placed.
- If Kai's data point is near the blue cluster, then we would predict that Kai belongs to class 0.
- If Kai's data point is near the orange cluster, then we would conclude that Kai belongs to class 1.
Kai's data point lies within the cluster of orange points. Since most of the nearby observations belong to class 1, we predict that Kai belongs to class 1 (female).
The reasoning above is an example of the nearest neighbor classification method. We predict the class of the new data point based on class of the nearest labeled data point. The idea is that data points that are close together are often more similar than data that are far apart. Therefore, if the nearest neighbor belongs to one class, then the new data is likely to be in the same class.
When a new data point clearly lies within a cluster of one class, as in the example above, then it's easy to determine which class the new data belongs to. But what if Kai's height is -0.1 standard units and weight is 1.3 standard units?
We plot Kai's data point in the scatterplot:
Now it is no longer clear which class Kai belongs to. Kai's data point lies between some blue and some orange data points, so we cannot decide the class simply by looking at the scatterplot.
To make a prediction, we need to determine which labeled data points are closest to Kai. We do this by calculating the distance between Kai and the other data points.
Find the Nearest Neighbor
From the scatterplot we see that the two nearest data points to Kai's data point are above and below Kai's data point, and both are slightly to the right. However, it is difficult to tell by eye which one is the nearest neighbor. We therefore measure the distance from Kai to each of these data points.
To calculate the distance, we first locate these data points in the dataset.
Index([3833, 9270], dtype='int64')
We found two indices for the two closest data points and save the two data points in two variables.
We now use the distance formula to find the distance between Kai's data point and closest_1 and find the distance between Kai's data point and closest_2.
In two-dimensional space, the distance between two points (x1, y1) and (x2, y2) is: $$ D = \sqrt{(x_1 - x_2)^2 + (y_1 - y_2)^2} $$
The formula is derived from the Pythagorean theorem, a foundational idea in geometry.
Although the formula looks complicated at first, we can break it down step by step in the an example below.
Distance between Point A and Point B: 5.0
Using the distance formula, we calculate the distance between closest_1 and Kai's data point and the distance between closest_2 and Kai's data point.
Distance to closest_1: 0.35
Distance to closest_2: 0.31
The nearest neighbor is closest_2. We can print the exact height and weight values for point 2.
closest_2 index: 9270
closest_2 height: -0.04
closest_2 weight: 1.26
We locate closest_2in the scatterplot using its height and weight. Since that data point is orange, we predict that Kai also belongs to class 1 (female).
K-Nearest Neighbors
In the example above, the nearest neighbor closest_2 is not the only data point that's close to Kai. The second nearest data point closest_1 is almost the same distance away. Rather than relying on just one nearby data point, we can make a more reliable prediction by considering several of the closest neighbors. This approach is called the k-nearest neighbors or kNN algorithm.
Instead of assigning Kai to the class of a single nearest neighbor, the k-nearest neighbors algorithm looks at the k closest labeled data points. The predicted class is the one that appears most often among these neighbors.
We display the scatterplot with Kai's data point again:
We notice that the closest neighbors to Kai's data point have height between -0.6 and 0.2 and weight between 1 and 2. We locate them in the sample dataset and show them in the plot.
Subset of rows based on locations closest to Kai:
| height | weight | gender | |
|---|---|---|---|
| 15716 | 0.198044 | 1.629668 | 0 |
| 3833 | -0.044330 | 1.879169 | 0 |
| 11902 | -0.529079 | 1.729469 | 1 |
| 9270 | -0.044330 | 1.255416 | 1 |
| 11659 | 0.198044 | 1.879169 | 0 |
We see that out of the five closest neighbors to Kai's data point, two are class 1 and three are class 0.
In KNN, the predicted class is the one that appears most often among these neighbors. Therefore, we predict that Kai belongs in class 0.
Using only one nearest neighbor, Kai is predicted to belong to class 1 because the closest observation belonged to class 1. However, when we consider the five nearest neighbors, 3/5 of the nearby observations belonged to class 0. As a result, the predicted class changed to class 0.
KNN makes predictions based on the majority vote of the nearest neighbors rather than relying on a single observation. This makes the prediction less sensitive to an individual data point that may not represent the surrounding neighborhood.
Choosing K
The value of k determines how many nearby data points are used to make a prediction.
- If k is too small, the prediction depends on only a few data points. Some unusual data points can have a large influence on the predicted result.
- If k is too large, the prediction includes many data points that may be too far from the new data point and may not be very similar to it. These more distant data points can outweigh the influence of the nearby data points.
In Kai's example above, using k=1 predicted class 1, while using k=5 predicted class 0. This shows that different values of k can lead to different predictions.
There is no single best value of k. The choice depends on the dataset and the problem being solved. For binary classification, where there are two categories, an odd value of k is often chosen to reduce the chance of a tie during the majority vote. And in practice, we often try several values of k and choose the one that gives the most accurate predictions.
Decision Boundary
So far we have used the k-nearest neighbor algorithm to predict the class of one new data point. For example, for Kai's data point, we predict that Kai belongs to class 0.
Now imagine many new data points at every location in the scatterplot, covering the entire plot. For each location we apply the k-nearest neighbor algorithm to predict a class, and we end up with predictions for all locations of the scatterplot. When we plot all the predictions for all locations on the plot, we will see that the plot is divided into regions, where all data points of the region belong to class 0 or all data points of the region belong to class 1.
The following code generates a plot showing the regions predicted to belong to class 0 and class 1, when k=5.
The code uses a Python machine learning library to generate the plot. Understanding this library is beyond the scope of this course, so you do not need to study the code. Instead, focus on understanding what the resulting plot shows.
In the plot above, we see both the scatterplot of height vs weight data and the predicted regions for each class. The orange regions represent predictions of class 1 and the blue regions represent predictions of class 0. The line separating each orange and blue region is the decision boundary. Any new data point on one side of the boundary is predicted to belong to class 0, while any new data point on the other side is predicted to belong to class 1.
A decision boundary is a useful way to visualize how a classification algorithm makes predictions. Instead of predicting the class of one new data point at a time, it shows the predicted class for every possible every possible combination of height and weight.
Automating the Nearest Neighbor Method
Up to now we've identified the nearest neighbors of a new data point by visually inspecting the scatterplot. This works for a small dataset, but it's not practical for larger datasets. Instead, we can use Python to automate the process by:
- calculating the distance between the new data point and every data point in the dataset
- sorting the distances from smallest to largest
- selecting the k data points with the shortest distances
First we write a function to calculate the distance between the new data point and one data point in the dataset
Next we write a function to go through each row of the sample dataset and find the distance between the new data point and the data point in the row. The function stores the distances in a list and return the list.
Now we use Kai's data point from the example above and find the 5 nearest neighbors.
500 distances calculated
k nearest neighbors to Kai's point:
| height | weight | gender | Distance | |
|---|---|---|---|---|
| 9270 | -0.044330 | 1.255416 | 1 | 0.312963 |
| 3833 | -0.044330 | 1.879169 | 0 | 0.345715 |
| 15716 | 0.198044 | 1.629668 | 0 | 0.357046 |
| 11902 | -0.529079 | 1.729469 | 1 | 0.419416 |
| 11659 | 0.198044 | 1.879169 | 0 | 0.479048 |
The calculated five nearest neighbors to Kai are the same five nearest neighbors that we identified earlier by visually inspecting the scatterplot.
We can also automate the prediction, since we know the five nearest neighbors.
Predicted gender for Kai: 0
We have now automated the k-nearest neighbors algorithm. Given any new data point, our code can calculate the distance to every data point in the dataset, identify the k nearest neighbors, and use their majority class to make a prediction.
Summary
In this chapter, we learned how the k-nearest neighbors (KNN) algorithm classifies new data by comparing it with nearby labeled data. We explored distance calculations, choosing the value of k, decision boundaries, and how to automate the process of finding the nearest neighbors and making predictions.
In the next chapter, we will use the same distance calculations and nearest-neighbor search to build a complete k-nearest neighbors classifier.


