k-nearest neighbours (kNN) is the machine-learning method with no training step. To predict the label of a new point, find the k stored examples closest to it and let them vote. For regression, average their values. The whole model is the training set plus a distance function. That makes kNN the easiest learner to explain and to update, and one of the easiest to get quietly wrong, because everything it does depends on what "close" means.
This article treats kNN as a learner. It covers the prediction rule in code, a worked example where a tie flips the answer, feature scaling, choosing k by cross-validation, what the theory guarantees, why high dimensions break it, and how to run it in production. The search structures that make queries fast have their own pages: k-d trees for low dimensions and HNSW for approximate search at scale.
The prediction rule, in code
Given training pairs (xi, yi), a query q, a distance d and an integer k, kNN runs in four steps.
- Compute d(q, xi) for every stored point, or ask an index for the k smallest.
- Keep the k nearest. This is the neighbourhood Nk(q).
- For classification, return the class with the most votes in Nk(q), optionally weighting each vote by 1/d. For regression, return the mean or distance-weighted mean of the neighbours' targets.
- For a probability, report the vote share of each class. With uniform weights the only possible values are 0, 1/k, 2/k and so on, up to 1.
The NumPy version below is complete and fast enough for tens of thousands of points. It uses the identity that the squared distance between q and x equals |q|2 - 2q.x + |x|2, so all distances come from one matrix product. It also stores the scaling statistics, which belong to the model just as much as the points do.
import numpy as np
def sq_dists(Q, X):
d = (Q**2).sum(1)[:, None] - 2 * Q @ X.T + (X**2).sum(1)[None, :]
return np.maximum(d, 0.0) # rounding can produce tiny negatives
class KNN:
def __init__(self, k=5, weighted=False):
self.k, self.weighted = k, weighted
def fit(self, X, y):
self.mu, self.sd = X.mean(0), X.std(0) + 1e-12
self.X = (X - self.mu) / self.sd # scaling is part of the model
self.y = np.asarray(y)
self.classes = np.unique(self.y)
return self
def neighbours(self, Q):
d = sq_dists((Q - self.mu) / self.sd, self.X)
idx = np.argpartition(d, self.k - 1, axis=1)[:, :self.k]
return idx, np.sqrt(np.take_along_axis(d, idx, axis=1))
def predict_proba(self, Q):
idx, d = self.neighbours(Q)
w = 1.0 / (d + 1e-9) if self.weighted else np.ones_like(d)
lab = self.y[idx]
votes = np.stack([(w * (lab == c)).sum(1) for c in self.classes], axis=1)
return votes / votes.sum(1, keepdims=True)
def predict(self, Q):
return self.classes[self.predict_proba(Q).argmax(1)]
def predict_value(self, Q): # kNN regression
idx, _ = self.neighbours(Q)
return self.y[idx].mean(1)Two details matter. argpartition finds the k smallest in linear time without fully sorting, but among equal distances it picks arbitrarily. And argmax breaks a tied vote in favour of whichever class sorts first. Both are tie-breaking policies, and the next section shows they can change the answer. The distance matrix has one row per query, so batch queries to keep it in memory.
Worked example: when a tie decides the answer
Take three class-A points at (1, 1), (2, 1) and (1, 2), three class-B points at (4, 4), (5, 4) and (4, 5), and a query at (3, 3). The Euclidean distances are:
| Point | Class | Distance to (3, 3) |
|---|---|---|
| (4, 4) | B | sqrt(2) = 1.414 |
| (2, 1) | A | sqrt(5) = 2.236 |
| (1, 2) | A | sqrt(5) = 2.236 |
| (5, 4) | B | sqrt(5) = 2.236 |
| (4, 5) | B | sqrt(5) = 2.236 |
| (1, 1) | A | sqrt(8) = 2.828 |
With k = 1 the answer is B. With k = 5 the answer is again B, by three votes to two. With k = 3, the nearest neighbour is B, and then four points tie at 2.236 for the remaining two places. If the training file lists the A points first and the tie is broken by file order, as a stable sort does, both A points get in and A wins two votes to one. Shuffle the training file and the prediction can change. Scikit-learn's documentation warns about exactly this: when neighbours k and k + 1 are equidistant with different labels, the result depends on the order of the training data.
Ties are common in real data, not a curiosity. Integer features, rounded measurements and duplicate rows all produce them. Decide the policy deliberately: include every point tied at the k-th distance, add a tiny deterministic jitter, or use distance weighting with a stable secondary key.
Distance is the model: scaling and metrics
kNN has no weights to learn, so the distance function is the model. Suppose one customer is age 30 with income 50,000. Another is age 60 with income 51,000, and a third is age 31 with income 53,000. On raw features, the 60-year-old is at distance sqrt(302 + 10002), about 1,000, while the 31-year-old is at about 3,000. The units of income dominate completely, and age is ignored.
Standardise each feature to zero mean and unit variance, or use robust scaling with the median and interquartile range when there are outliers. Fit the scaler on training data only and apply it unchanged at query time. Beyond scaling, the metric is a modelling choice:
- Euclidean for dense, comparably scaled numeric features.
- Cosine for embeddings and text vectors, where direction matters more than length. On unit-normalised vectors, cosine and Euclidean rank neighbours identically.
- Manhattan when you want less influence from a single large coordinate difference.
- Hamming or Jaccard for binary and set-valued features. One-hot encoded categories mixed into a Euclidean distance quietly weight each category change as sqrt(2).
- Learned metrics such as an embedding model or metric learning, when raw features do not reflect similarity. For images, text and users, this is usually the case.
Choosing k
Small k gives a jagged boundary that follows every noisy point: low bias, high variance. With k = 1 the training error is zero, which tells you nothing. Large k smooths toward the majority class: at k = n every query gets the overall most common label. Choose k by cross-validation, not by a rule of thumb. For binary problems, odd k avoids tied votes, although tied distances can still occur. Put the scaler inside the pipeline so it is refitted on each fold.
from sklearn.pipeline import make_pipeline
from sklearn.preprocessing import StandardScaler
from sklearn.neighbors import KNeighborsClassifier
from sklearn.model_selection import GridSearchCV
pipe = make_pipeline(StandardScaler(), KNeighborsClassifier())
grid = GridSearchCV(
pipe,
{"kneighborsclassifier__n_neighbors": [1, 3, 5, 9, 15, 25, 41],
"kneighborsclassifier__weights": ["uniform", "distance"]},
cv=5, scoring="balanced_accuracy")
grid.fit(X_train, y_train)
print(grid.best_params_, grid.best_score_)Use a balanced metric when classes are skewed, because a large k drifts toward always predicting the majority. Plot the validation score against k. A broad, flat optimum is reassuring, while a sharp spike at k = 1 often means near-duplicates are leaking between folds.
What the theory guarantees
kNN comes with unusually clean guarantees. Cover and Hart (1967) showed that as the training set grows without limit, the error of the 1-nearest-neighbour rule is at most twice the Bayes error, the lowest error any classifier can reach. In that sense, half of the information needed for classification sits in the nearest neighbour. Stone (1977) showed that kNN is universally consistent when k grows with n but more slowly, so k tends to infinity while k/n tends to zero. It then converges to the Bayes error for any distribution.
These are asymptotic statements. How fast kNN gets there depends on dimension, and that is where it struggles.
The curse of dimensionality
Spread points uniformly in a d-dimensional unit cube and ask for a sub-cube that holds 1% of them. Its edge length is 0.011/d. In 10 dimensions that is 0.631, and in 100 dimensions it is 0.955. A neighbourhood that holds 1% of the data spans most of the range of every feature, so it is not local in any useful sense. A related effect, studied by Beyer and colleagues (1999), is that for many distributions the ratio of the farthest to the nearest distance tends to 1 as d grows. Neighbours stop being meaningfully nearer than anything else.
Real data is rarely uniform. It usually lies near a lower-dimensional surface, which is why kNN on good embeddings works well. The practical rules follow from that: reduce or learn the representation first (PCA, a trained encoder), drop irrelevant features, which add noise to every distance, and check neighbour quality by eye on a sample of queries.
Making queries fast
Brute force costs O(nd) per query. With batched queries it becomes a matrix product, so it runs at BLAS or GPU speed and is exact. For moderate n, it is often the fastest option in practice. Past that, choose an index:
| Index | Exact? | Fits | Watch out for |
|---|---|---|---|
| Brute force (GEMM) | Yes | Up to moderate n; any d; GPUs | Linear cost per query as n grows |
| k-d tree | Yes | Low d (roughly tens at most) | Degrades toward brute force as d rises |
| Ball tree | Yes | Moderate d, non-Euclidean metrics | Build cost; still hit by high d |
| HNSW, IVF and other ANN | No (recall below 100%) | Large n, high-dimensional embeddings | Recall must be measured; memory overhead |
Approximate search is often fine for kNN classification. Swapping one of k neighbours for a slightly farther one rarely changes a majority vote. Measure it anyway: compare predictions from the exact and approximate index on a held-out sample before you switch.
Operational guidance
- Memory is the model. Ten million 768-dimensional float32 vectors take 30.7 GB before any index overhead. Budget for it, or condense the set (keep only points near the decision boundary) or quantise vectors.
- Updates are cheap. Adding or deleting a training example is an append or a delete, which suits fast-changing catalogues and data-deletion requests. Graph indexes such as HNSW often delete by marking entries, so plan periodic rebuilds.
- Explanations come free. Return the neighbours with each prediction. Reviewers can see why a decision was made and spot label errors.
- Probabilities are coarse. Vote shares come in steps of 1/k and are poorly calibrated. Recalibrate them, as described in model calibration, before using them for thresholds or ranking.
- Monitor drift. Track the median distance to the k-th neighbour for live queries. A rising value means traffic is moving into regions with no training data.
Failure modes
- Unscaled features, where one large-unit column decides every distance.
- Duplicate leakage. Near-duplicates across the train and test split make 1-NN look excellent offline and disappoint in production. De-duplicate, or split by entity.
- Class imbalance. The majority class floods every neighbourhood. Use distance weighting, balanced scoring and per-class thresholds.
- Order-dependent ties, as in the worked example. Predictions change when data is reloaded in a different order.
- Querying a point against itself during evaluation, which returns the point as its own neighbour at distance 0.
Trade-offs against other classifiers
| kNN | Logistic regression | Naive Bayes | Kernel SVM | |
|---|---|---|---|---|
| Training | None (store data) | Convex optimisation | Count statistics | Quadratic program |
| Prediction cost | Grows with n | O(d) | O(d) | Grows with support vectors |
| Boundary | Any shape, local | Linear in features | Linear or simple | Any shape via kernel |
| Needs scaling | Critical | For optimisation and penalties | No | Yes |
kNN is a strong baseline when the representation is good and n is manageable, and a poor choice for many raw, irrelevant features. Compare it with naive Bayes and kernel SVMs on the same split.
What to do next
- Implement the KNN class above and reproduce the worked example, including the k = 3 flip when you reverse the training order.
- Put a scaler and the classifier in one pipeline and grid-search k and weighting with a balanced metric.
- Plot validation score against k, and check for duplicate leakage if k = 1 wins.
- Measure the median k-th-neighbour distance on validation data and set a drift alert on the live value.
- If n or d is large, benchmark brute-force GEMM against an index, and measure ANN recall and prediction agreement before switching.
- Calibrate probabilities if anything downstream thresholds them.