Skip to main content
Workforce LibreTexts

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.

    import pandas as pd
    import numpy as np
    import matplotlib.pyplot as plt
    
    url = 'https://raw.githubusercontent.com/DeAnzaDataScience/CIS11/refs/heads/main/datasets_notes/cdc.csv'
    health = pd.read_csv(url)
    display(health.head())
    print("Number of rows:", health.shape[0])
    Output?
      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.

    data = health[['height', 'weight', 'gender']].copy()
    data.head()
    Output?
      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.

    # the DataFrame has a map method that can be used to map values in a column to new values
    
    data['gender'] = data['gender'].map({'m': 0, 'f': 1})
    data.head()
    data.head()
    Output?
      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.

    print("Maximum height:", data['height'].max())
    print("Minimum height:", data['height'].min())
    print("Maximum weight:", data['weight'].max())
    print("Minimum weight:", data['weight'].min())
    Output?
    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.

    def standard_units(array):
        return (array - np.mean(array))/np.std(array)
    
    data['height'] = standard_units(data['height'])
    data['weight'] = standard_units(data['weight'])
    data.head()
    Output?
      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.

    plt.figure(figsize=(4,4))
    plt.scatter(data['height'], data['weight'], alpha=0.5)
    plt.xlabel('Height')
    plt.ylabel('Weight')
    plt.grid()
    plt.show()
    Output?

    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.

    # create a sample of 500 rows from data
    sample = data.sample(500, random_state=1)
    
    # plot a scatterplot of the sample data
    plt.figure(figsize=(4,4))
    plt.scatter(sample['height'], sample['weight'], alpha=0.5)
    plt.xlabel('Height')
    plt.ylabel('Weight')
    plt.grid()
    plt.show()
    Output?

    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.

    # group by gender: the result will be separated into two groups
    groups = sample.groupby('gender')
    
    plt.figure(figsize=(4,4))
    for gender, group in groups:
        plt.scatter(group['height'], group['weight'], alpha=0.4, label=gender)
    
    plt.xlabel('Height')
    plt.ylabel('Weight')
    plt.legend()
    plt.grid()
    plt.show()
    Output?

    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_height = -1.2
    kai_weight = -0.5
    
    plt.figure(figsize=(4,4))
    for gender, group in groups:
        plt.scatter(group['height'], group['weight'], alpha=0.4, label=gender)
    plt.scatter(kai_height, kai_weight, color='firebrick', s=50, label='Kai')
    plt.xlabel('Height')
    plt.ylabel('Weight')
    plt.legend(fontsize=9)
    plt.grid()
    plt.show()
    Output?

    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:

    kai_height = -0.15
    kai_weight = 1.55
    
    plt.figure(figsize=(4,4))
    for gender, group in groups:
        plt.scatter(group['height'], group['weight'], alpha=0.4, label=gender)
    plt.scatter(kai_height, kai_weight, color='firebrick', s=50, label='Kai')
    plt.xlabel('Height')
    plt.ylabel('Weight')
    plt.legend(fontsize=9)
    plt.grid()
    plt.show()
    Output?

    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.

    # earch for the two data points with height between -0.5 and 0 and weight between 1 and 2
    condition = ((sample['height'] >= -0.5) & (sample['height'] <= 0) &
                 (sample['weight'] >= 1) & (sample['weight'] <= 2))
    
    # find the index of the closest data points to Kai's data point
    closest_data_index = sample.loc[condition].index
    closest_data_index
    Output?
    Index([3833, 9270], dtype='int64')

    We found two indices for the two closest data points and save the two data points in two variables.

    closest_1 = sample.loc[closest_data_index[0]]
    closest_2 = sample.loc[closest_data_index[1]]
    Output?

    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.

    # Illustration of the distance between two points in a 2D space
    
    # create 2 sample points
    point_A = (0, 0)   # (x1, y1)
    point_B = (3, 4)   # (x2, y2)
    
    # Calculate the rise and run
    # - the rise is the difference along the y-axis between the points
    # - the run is the difference along the x-axis between the points
    run = point_A[0] - point_B[0]    # x1 - x2
    rise = point_A[1] - point_B[1]   # y1 - y2
    
    # Calculate the distance
    distance = np.sqrt(run**2 + rise**2)
    print("Distance between Point A and Point B:", distance)
    
    # Plot
    plt.figure(figsize=(3,3))
    
    plt.plot(point_A[0], point_A[1], 'o', color='blue', label="Point A")
    plt.plot(point_B[0], point_B[1], 'o', color='orange', label="Point B")
    
    plt.plot([point_A[0], point_B[0]], [point_A[1], point_A[1]], color='brown',
             linestyle='--', label="Run")
    plt.plot([point_B[0], point_B[0]], [point_A[1], point_B[1]], color='green',
             linestyle='--', label="Rise")
    plt.plot([point_A[0], point_B[0]], [point_A[1], point_B[1]], color='red',
             linestyle='-', label="Distance")
    
    plt.legend(bbox_to_anchor=(1.05, 1), loc='upper left')
    plt.xlabel('x')
    plt.ylabel('y')
    plt.grid()
    plt.show()
    Output?
    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_1 = np.sqrt((closest_1['height'] - kai_height)**2 + (closest_1['weight'] - kai_weight)**2)
    distance_2 = np.sqrt((closest_2['height'] - kai_height)**2 + (closest_2['weight'] - kai_weight)**2)
    print("Distance to closest_1:", round(distance_1, 2))
    print("Distance to closest_2:", round(distance_2, 2))
    Output?
    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.

    # find the index of closest_2 and print height and weight using that index
    idx = closest_2.name
    print("closest_2 index:", idx)
    print("closest_2 height:", round(data.at[idx, 'height'],2))
    print("closest_2 weight:", round(data.at[idx, 'weight'],2))
    Output?
    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:

    kai_height = -0.15
    kai_weight = 1.55
    
    plt.figure(figsize=(4,4))
    for gender, group in groups:
        plt.scatter(group['height'], group['weight'], alpha=0.4, label=gender)
    plt.scatter(kai_height, kai_weight, color='firebrick', s=50, label='Kai')
    plt.xlabel('Height')
    plt.ylabel('Weight')
    plt.legend(fontsize=9)
    plt.grid()
    plt.show()
    Output?

    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.

    condition = ((sample['height'] > -0.6) & (sample['height'] < 0.2) &
                 (sample['weight'] > 1.25) & (sample['weight'] < 2))
    
    subset_rows = sample.loc[condition]
    print("Subset of rows based on locations closest to Kai:")
    subset_rows
    Output?
    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
    plt.figure(figsize=(4,4))
    for gender, group in groups:
        plt.scatter(group['height'], group['weight'], alpha=0.5, label=gender)
    plt.scatter(kai_height, kai_weight, color='firebrick', s=50, label='Kai')
    plt.scatter(subset_rows['height'], subset_rows['weight'], color = 'cornsilk', alpha=0.3, linewidth=3, edgecolors='black', label='Closest')
    plt.xlabel('Height')
    plt.ylabel('Weight')
    plt.legend(fontsize=8)
    plt.grid()
    plt.show()
    Output?

    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.

    # create the features and target variable for classification
    X = sample[['height', 'weight']]   # input features
    y = sample['gender']               # label or target variable
    
    # use Python machine learning modules to create a classification model
    from sklearn.neighbors import KNeighborsClassifier
    from matplotlib.colors import ListedColormap
    # the model is created with the number of neighbors k set to 5
    knn = KNeighborsClassifier(n_neighbors=5)
    # the model is given the input features and the target variable 
    # so it can learn the relationship between the features and the target variable
    # and be able to calculate the distance between new data points and the input features
    knn.fit(X, y)
    
    # create a mesh of new data points over the range of the input features
    x_min, x_max = X['height'].min() - 0.5, X['height'].max() + 0.5
    y_min, y_max = X['weight'].min() - 0.5, X['weight'].max() + 0.5
    
    xx, yy = np.meshgrid( np.arange(x_min, x_max, 0.02),
                          np.arange(y_min, y_max, 0.02))
    
    # the model predicts the class for each new data point
    grid = pd.DataFrame({'height': xx.ravel(),
                         'weight': yy.ravel()})
    Z = knn.predict(grid)
    Z = Z.reshape(xx.shape)
    
    # plot the resulting classification regions and the input features
    cmap = ListedColormap(['#1f77b4', '#ff7f0e'])
    plt.contourf(xx, yy, Z, alpha=0.3, cmap=cmap)
    plt.scatter(X['height'], X['weight'], c=y, cmap=cmap, alpha=0.6)
    plt.xlabel("Height")
    plt.ylabel("Weight")
    plt.grid()
    plt.show()
    Output?

    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

    # calculate the distance between a new data point and one data point in the dataset
    # new_point is a new data point with x and y values (height and weight)
    # row is a row or one data point in the dataset with x and y values (height and weight)
    def distance_from_point(new_point, row):
      # convert each data point into an array
      new_point = np.array(new_point)   # x1, y1
      row = np.array(row)               # x2, y2
      # subtract the 2 arrays
      difference = new_point - row      # x1 - x2, y1 - y2
      # square each difference
      squared_difference = difference ** 2      # (x1 - x2)^2, (y1 - y2)^2
      # add the squared differences and take the square root
      D = np.sqrt(np.sum(squared_difference))   
      return D
    Output?

    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.

    def find_distances(new_point, dataset):
        # create a list to store the distances
        distances = []
        for row in dataset: 
            # calculate the distance between the new point and the current point
            a_distance = distance_from_point(new_point, row)
            distances.append(a_distance)
        return distances
    Output?

    Now we use Kai's data point from the example above and find the 5 nearest neighbors.

    kai_height = -0.15  # same location as before
    kai_weight = 1.55
    kai_point = [kai_height, kai_weight]
    
    # create a DataFrame with the features (height and weight) from sample
    features = sample[['height', 'weight']] 
    # find all the distances
    distances = find_distances(kai_point, features.values)
    print(len(distances), "distances calculated")
    Output?
    500 distances calculated
    
    results = sample.copy()
    results["Distance"] = distances
    
    k = 5
    k_nearest_neighbors = results.sort_values(by='Distance').head(k)
    print("k nearest neighbors to Kai's point:")
    k_nearest_neighbors
    Output?
    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.

    # function to find the majority class among the k nearest neighbors
    def majority(k_nearest):
        # find the count of 1's
        ones = np.sum(k_nearest.gender == 1)
        # find the count of 0's
        zeros = np.sum(k_nearest.gender == 0)
        # return the larger of the two
        if ones > zeros:
            return 1
        else:
            return 0
        
    print("Predicted gender for Kai:", majority(k_nearest_neighbors))
    Output?
    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.


      This page titled 15: Classification was last modified on Fri, 25 Sep 2026 01:22:23 GMT and is shared under a CC BY 4.0 license and was authored, remixed, and/or curated by Clare Nguyen.

      • Was this article helpful?