第32章 Top-K和KNN
给一个查询向量,遍历数据库里的记录,算距离,保留距离最小的 K 条。这就是 KNN(K-Nearest Neighbors)。它的基本动作不复杂:算距离、比大小、留 K 个。
32.1 问题从哪来#
上一章学会了算两个向量之间的距离。只看学习时长和出勤率这两个维度时,Alice 和 Bob 的距离约为 16.00,Alice 和 Carol 的距离约为 6.00,一比就知道谁更像。
但数据库里不会只有两条记录。真实的问题是这样的:
数据库里有 1000 个学生。我输入一个新学生的特征向量,想找数据库里和他最像的 5 个人。
这就是 KNN(K-Nearest Neighbors,K 近邻) 问题:给一个查询点,从所有记录里找出距离最近的 K 条。
这一章里的 Top-K 指的是"距离最小的 K 个"。如果把每条记录到查询点的距离都算出来,Top-K 就是距离表里排在最前面的 K 条。
32.2 先看一个例子#
假设数据库里有 6 个学生,每个学生用二维向量描述(学习时长、出勤率):
| id | name | vec |
|---|---|---|
| 1 | Alice | [21.0, 0.95] |
| 2 | Bob | [5.0, 0.60] |
| 3 | Carol | [15.0, 0.85] |
| 4 | Dave | [18.0, 0.90] |
| 5 | Eve | [8.0, 0.65] |
| 6 | Frank | [23.0, 0.92] |
现在有一个查询点 Q = [16.0, 0.88],想找和 Q 最像的 3 个学生(K = 3)。
先算 Q 到每个点的欧氏距离:
| name | 距离 |
|---|---|
| Alice | 5.00 |
| Bob | 11.00 |
| Carol | 1.00 |
| Dave | 2.00 |
| Eve | 8.00 |
| Frank | 7.00 |
排序后取前 3 个:Carol(1.00)、Dave(2.00)、Alice(5.00)。
这就是 KNN 的全部逻辑:算距离、排序、取前 K 个。
32.3 最小实验#
基本三步:
- 遍历数据库所有记录
- 计算查询向量和每条记录的距离
- 保留距离最小的 K 条记录
第 3 步有两种做法:
- 做法 A:先算完所有距离,存进数组,排序,取前 K。简单直接。
- 做法 B:边扫描边维护一个"当前最近 K 个"的列表。记录多的时候更省空间。
做法 A 适合记录少的场景,代码最容易理解。做法 B 更接近实际使用中的 Top-K 写法。本章把两种方法都写一遍。
32.4 做法 A:全部算完再排序#
思路最直接:算出每条记录到查询点的距离,存进一个数组,排序,取前 K 个。
这段演示里有一个很小的 StudentDB,它只是用来放几条测试数据。读代码时重点看 Result all[]、euclidean_dist、qsort 和最后取前 K 个的循环。
#include <stdio.h>
#include <math.h>
#include <stdlib.h>
#include <string.h>
#define DIM 2
#define MAX_STUDENTS 100
#define NAME_LEN 32
struct Student {
int id;
char name[NAME_LEN];
float vec[DIM];
};
struct StudentDB {
struct Student rows[MAX_STUDENTS];
int count;
};
struct Result {
int index; // 在数据库中的下标
float dist; // 到查询点的距离
};
void db_init(struct StudentDB *db)
{
// 将记录数置零,数据库初始化为空
db->count = 0;
}
int db_insert(struct StudentDB *db, int id, const char *name, const float vec[DIM])
{
// 数据库已满则插入失败
if (db->count >= MAX_STUDENTS) return 0;
// 取末尾位置写入新记录
struct Student *s = &db->rows[db->count];
s->id = id;
snprintf(s->name, sizeof(s->name), "%s", name);
// 逐维度复制向量分量
for (int i = 0; i < DIM; i++) s->vec[i] = vec[i];
// 记录数加 1
db->count++;
return 1;
}
float euclidean_dist(const float a[DIM], const float b[DIM])
{
float sum = 0.0f;
for (int i = 0; i < DIM; i++) {
// 各维度差值的平方累加
float d = a[i] - b[i];
sum += d * d;
}
// 开平方得到欧氏距离
return sqrtf(sum);
}
// 用于 qsort:按距离从小到大
int cmp_by_dist(const void *a, const void *b)
{
float da = ((const struct Result *)a)->dist;
float db = ((const struct Result *)b)->dist;
if (da < db) return -1;
if (da > db) return 1;
return 0;
}
int knn_sort(const struct StudentDB *db, const float query[DIM],
int k, struct Result *out)
{
// 存放所有记录的(下标, 距离)
struct Result all[MAX_STUDENTS];
int n = db->count;
// 参数无效时直接返回 0
if (k <= 0 || n <= 0) {
return 0;
}
// 遍历数据库,计算每条记录到查询点的距离
for (int i = 0; i < n; i++) {
all[i].index = i;
all[i].dist = euclidean_dist(query, db->rows[i].vec);
}
// 按距离从小到大排序
qsort(all, n, sizeof(struct Result), cmp_by_dist);
// 取前 K 个(记录不足时取全部)
int found = (k < n) ? k : n;
for (int i = 0; i < found; i++) {
out[i] = all[i];
}
return found;
}
int main(void)
{
struct StudentDB db;
// 初始化数据库
db_init(&db);
// 插入 6 条学生记录
float v1[] = {21.0f, 0.95f}; db_insert(&db, 1, "Alice", v1);
float v2[] = {5.0f, 0.60f}; db_insert(&db, 2, "Bob", v2);
float v3[] = {15.0f, 0.85f}; db_insert(&db, 3, "Carol", v3);
float v4[] = {18.0f, 0.90f}; db_insert(&db, 4, "Dave", v4);
float v5[] = {8.0f, 0.65f}; db_insert(&db, 5, "Eve", v5);
float v6[] = {23.0f, 0.92f}; db_insert(&db, 6, "Frank", v6);
// 查询向量和 K 值
float query[DIM] = {16.0f, 0.88f};
int k = 3;
// 执行 KNN 查询
struct Result results[MAX_STUDENTS];
int found = knn_sort(&db, query, k, results);
// 打印结果
printf("Query point: [%.2f, %.2f]\n", query[0], query[1]);
printf("Top-%d nearest neighbors:\n", k);
for (int i = 0; i < found; i++) {
const struct Student *s = &db.rows[results[i].index];
printf(" %d. %-8s dist=%.2f\n", i + 1, s->name, results[i].dist);
}
return 0;
}32.5 编译运行#
保存为 knn_demo.c。这段代码用了 sqrtf,编译时要链接数学库,所以命令里有 -lm:
$ gcc -std=c11 -Wall -Wextra knn_demo.c -o knn_demo.exe -lm
$ .\knn_demo.exe
运行结果:
Query point: [16.00, 0.88]
Top-3 nearest neighbors:
1. Carol dist=1.00
2. Dave dist=2.00
3. Alice dist=5.00
Carol、Dave、Alice 距离查询点最近。Bob 和 Eve 离得远,Frank 也不近。
32.6 做法 B:边扫描边保留 Top-K#
做法 A 的问题:如果数据库有 100 万条记录,就要开一个 100 万元素的数组来存所有距离。qsort 本身也要花时间。
做法 B 的思路不同:维护一个长度为 K 的"当前最佳"列表。每扫描一条新记录,算出距离,如果比列表里最远的那个还近,就替换掉最远的。
列表初始时是空的。前 K 条记录直接填进去。从第 K + 1 条开始,每条都要和列表里距离最大的那个比。
// 找 results 里距离最大的那个的下标
int find_max_index(const struct Result *results, int count)
{
int max_i = 0;
for (int i = 1; i < count; i++) {
if (results[i].dist > results[max_i].dist) {
max_i = i;
}
}
return max_i;
}
int knn_stream(const struct StudentDB *db, const float query[DIM],
int k, struct Result *out)
{
int n = db->count;
// 参数无效时直接返回 0
if (k <= 0 || n <= 0) {
return 0;
}
int cap = (k < n) ? k : n; // 实际能取的数量
int filled = 0;
// 逐条扫描数据库
for (int i = 0; i < n; i++) {
float d = euclidean_dist(query, db->rows[i].vec);
if (filled < cap) {
// 还没填满,直接加进去
out[filled].index = i;
out[filled].dist = d;
filled++;
} else {
// 已满,看是否比最远的更近
int worst = find_max_index(out, cap);
if (d < out[worst].dist) {
// 替换掉当前最远的那条
out[worst].index = i;
out[worst].dist = d;
}
}
}
return filled;
}调用方式和做法 A 一样,只换一个函数名:
int found = knn_stream(&db, query, k, results);保留下来的 Top-K 集合和做法 A 一样。不过 knn_stream 只是维护一个候选列表,列表里的元素不一定已经按距离排好。如果要按距离从小到大打印,可以再对 results[0..found-1] 排一次序。
32.7 Top-K 列表怎么更新#
做法 B 的核心是那个"替换"操作。用 6 条记录、K = 3 走一遍过程:
| 扫描到 | 距离 | 列表状态 | 操作 |
|---|---|---|---|
| Alice | 5.00 | [Alice] | 未满,直接加入 |
| Bob | 11.00 | [Alice, Bob] | 未满,直接加入 |
| Carol | 1.00 | [Alice, Bob, Carol] | 未满,直接加入 |
| Dave | 2.00 | [Alice, Dave, Carol] | 已满,最远的是 Bob(11.00),2.00 < 11.00,替换 Bob |
| Eve | 8.00 | [Alice, Dave, Carol] | 已满,最远的是 Alice(5.00),8.00 > 5.00,不替换 |
| Frank | 7.00 | [Alice, Dave, Carol] | 已满,最远的是 Alice(5.00),7.00 > 5.00,不替换 |
最终保留下来的 3 条记录是 Alice、Dave、Carol。把它们按距离排一下,就是 Carol、Dave、Alice。
前三条记录先把候选列表填满:
列表满了以后,只和当前最远的候选比较:
如果数据库里再加一条 Grace,她的向量是 [17.0, 0.89],距离查询点约等于 1.00,和 Carol 很接近。扫描到 Grace 时,列表已满,最远的是 Alice(5.00),1.00 < 5.00,替换掉 Alice。此时列表里保留 Grace、Dave、Carol;如果再按距离排序,结果是 Grace、Carol、Dave。
32.8 数据/内存/流程里发生了什么#
32.8.1 两种做法的内存对比#
| 做法 A(排序) | 做法 B(流式) | |
|---|---|---|
| 距离数组大小 | n 个 | K 个 |
| 维护开销 | 排序 O(n log n) | 线性维护 O(nK) |
| 额外空间 | O(n) | O(K) |
做法 A 需要开一个和记录数一样大的数组。做法 B 只需要 K 个空间。当 n = 100 万、K = 5 时,做法 A 需要 100 万个 struct Result,做法 B 只需要 5 个。
32.8.2 全量扫描的成本#
无论哪种做法,都要遍历每一条记录,算一次距离。这就是全量扫描(full scan)。
| 记录数 | 距离计算次数 | 每次计算耗时 | 总耗时(估算) |
|---|---|---|---|
| 100 | 100 | ~1 μs | ~0.1 ms |
| 10,000 | 10,000 | ~1 μs | ~10 ms |
| 100 万 | 1,000,000 | ~1 μs | ~1 s |
记录少的时候感觉不到。记录到了百万级,每次查询都要算 100 万次距离,延迟就明显了。
这是 KNN 最根本的限制:没有索引加速的情况下,只能一条一条算。在低维数据里,KD-tree 之类的数据结构常常能减少比较次数;但最坏情况下仍可能退化成 O(n),维度高了效果也会下降。
32.8.3 做法 B 的替换操作#
find_max_index 每次都要遍历 K 个元素找最大值。K 很小时(比如 5)这个开销可以忽略。K 很大时,可以用最大堆(max-heap)来加速——堆顶永远是最大值,每扫描一条记录时,维护候选集的开销可以降到 O(log K)。
对于这个教材的小数据库,线性扫描 K 个元素足够了。
32.9 把 KNN 接进数据库#
KNN 不是一个独立程序,它是数据库的一种查询方式。把 knn_stream 封装成数据库的一个操作:
void db_knn(const struct StudentDB *db, const float query[DIM], int k)
{
// 用流式算法找出 Top-K
struct Result results[MAX_STUDENTS];
int found = knn_stream(db, query, k, results);
// 结果按距离从小到大排序
qsort(results, found, sizeof(results[0]), cmp_by_dist);
// 打印排序后的结果
printf("Query point: [%.2f, %.2f] Top-%d:\n", query[0], query[1], k);
for (int i = 0; i < found; i++) {
const struct Student *s = &db->rows[results[i].index];
printf(" %d. id=%d %-8s dist=%.2f\n",
i + 1, s->id, s->name, results[i].dist);
}
}调用方只需要一行:
db_knn(&db, query, 5);数据库从"只能查 id = 3"变成了"查和这个向量最像的 5 条记录"。
32.10 常见坑#
坑 1:K 大于记录数。 数据库只有 3 条记录,但 k = 5。代码里 cap = (k < n) ? k : n 做了保护,输出只有 3 条。这不是 bug,是数据不够。
坑 2:距离算错。 euclidean_dist 里忘了取平方根,返回的是距离的平方。对欧氏距离这种非负距离来说,比较距离平方和比较距离本身会得到相同顺序,Top-K 不会错。但打印距离值给人看时数字会很奇怪。
坑 3:查询向量维度不对。 数据库里的向量是 2 维,查询向量也应该是 2 维。本章代码用 DIM 控制循环,euclidean_dist 只会读取 DIM 个分量;如果查询向量实际给少了,就可能读到不该读的内存。
坑 4:做法 A 的比较函数写反。 如果 cmp_by_dist 返回相反的值,排序变成从大到小,取前 K 个就变成了最远的 K 个——最不像的 K 个。
坑 5:做法 B 没处理 K 的边界。 如果 k <= 0,应该直接返回 0;如果 n < K,要用 cap = (k < n) ? k : n。这些保护能避免函数去访问没有意义的候选位置。
32.11 自己试试看#
Q1:改 K 值。 把 K 从 3 改成 1、改成 6,观察输出变化。K = 1 时就是"最近的那个"。
提示:把 k 改成不同值后重新编译运行。k = 1 时只输出一条;k 大于记录总数时,代码里的 cap 保护会让输出等于记录总数,不会越界。
Q2:加更多学生。 往数据库里加 10 条、20 条记录,看 Top-3 会不会变。
提示:在数据库初始化代码里多加几个学生,保持查询向量不变。原来排前三的记录可能被新数据挤掉——新数据离查询点更近,就取代了原来的位置。
Q3:用三维向量。 把 DIM 改成 3,给每个学生加一个维度,查询向量也改成 3 维,验证距离计算是否正确。
提示:统一用 DIM 宏控制所有循环和数组大小。查询向量声明为 float query[DIM],和数据库向量维度一致。打印结果时注意 printf 里维度相关的地方也要跟着变。
Q4:打印全部距离。 在 knn_stream 里加一行 printf,每扫描一条记录就打印它的距离。观察哪些被保留、哪些被替换。
提示:在 knn_stream 的 for 循环里,计算完 d 之后加 printf(" %s: %.2f\n", db->rows[i].name, d);。观察输出:距离小的记录排在前面,距离大的可能被后面的更小距离替换掉。
Q5:两种做法结果对比。 用同一组数据分别调用 knn_sort 和 knn_stream,打印结果对比。Top-K 集合应该一样。
提示:先调 knn_sort 存一份结果数组,再调 knn_stream 存另一份。用循环对比两组的 id,如果顺序不同但集合相同(都包含同样的 K 条记录),说明排序稳定性导致了顺序差异,但结果仍然正确。
拓展阅读#
本章的流式 Top-K 每次用 find_max_index 找当前候选里最远的那条。K 很小时这样写很清楚;K 变大时,可以用最大堆保存当前 K 个候选。最大堆的堆顶放“当前最差的候选”,新记录只需要和堆顶比较,再决定是否替换。
堆也常用来实现优先队列。优先队列按优先级出队,每次取出的都是优先级最高或最低的元素,和进入顺序无关。向量搜索里还会见到 KD-tree,它把低维空间一层一层切开,减少需要比较的点。维度升高后,一次查询要访问的节点变多,KD-tree 相对逐个比较的优势会缩小。
下一章的问题#
KNN 解决了一个具体问题:给一个查询点,找数据库里最像的 K 条记录。
但 KNN 需要一个前提:先给出查询点。如果数据库里有 1000 个学生,没有任何标签,也没有查询点,能不能让程序只看这些数据本身,自动把学生分成几组?
这个问题不再是"谁离查询点最近",而是"哪些点自然聚在一起"。这就进入了 KMeans 聚类 要解决的场景。