第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 个学生,每个学生用二维向量描述(学习时长、出勤率):

idnamevec
1Alice[21.0, 0.95]
2Bob[5.0, 0.60]
3Carol[15.0, 0.85]
4Dave[18.0, 0.90]
5Eve[8.0, 0.65]
6Frank[23.0, 0.92]

现在有一个查询点 Q = [16.0, 0.88],想找和 Q 最像的 3 个学生(K = 3)。

查询点 Q 和数据库里的 6 个点分布在二维平面上

先算 Q 到每个点的欧氏距离:

name距离
Alice5.00
Bob11.00
Carol1.00
Dave2.00
Eve8.00
Frank7.00

排序后取前 3 个:Carol(1.00)、Dave(2.00)、Alice(5.00)。

距离最近的 3 个点被标记出来,其余点用灰色表示

这就是 KNN 的全部逻辑:算距离、排序、取前 K 个。


32.3 最小实验#

基本三步:

  1. 遍历数据库所有记录
  2. 计算查询向量和每条记录的距离
  3. 保留距离最小的 K 条记录

第 3 步有两种做法:

  • 做法 A:先算完所有距离,存进数组,排序,取前 K。简单直接。
  • 做法 B:边扫描边维护一个"当前最近 K 个"的列表。记录多的时候更省空间。

做法 A 适合记录少的场景,代码最容易理解。做法 B 更接近实际使用中的 Top-K 写法。本章把两种方法都写一遍。


32.4 做法 A:全部算完再排序#

思路最直接:算出每条记录到查询点的距离,存进一个数组,排序,取前 K 个。

这段演示里有一个很小的 StudentDB,它只是用来放几条测试数据。读代码时重点看 Result all[]euclidean_distqsort 和最后取前 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 走一遍过程:

扫描到距离列表状态操作
Alice5.00[Alice]未满,直接加入
Bob11.00[Alice, Bob]未满,直接加入
Carol1.00[Alice, Bob, Carol]未满,直接加入
Dave2.00[Alice, Dave, Carol]已满,最远的是 Bob(11.00),2.00 < 11.00,替换 Bob
Eve8.00[Alice, Dave, Carol]已满,最远的是 Alice(5.00),8.00 > 5.00,不替换
Frank7.00[Alice, Dave, Carol]已满,最远的是 Alice(5.00),7.00 > 5.00,不替换

最终保留下来的 3 条记录是 Alice、Dave、Carol。把它们按距离排一下,就是 Carol、Dave、Alice。

前三条记录先把候选列表填满:

Top-K 列表未满时,前三条记录直接加入候选列表

列表满了以后,只和当前最远的候选比较:

Top-K 列表已满时,只有更近的记录才会替换当前最远候选

如果数据库里再加一条 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)。

全量扫描:查询点和数据库里每条记录都要算一次距离

记录数距离计算次数每次计算耗时总耗时(估算)
100100~1 μs~0.1 ms
10,00010,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_streamfor 循环里,计算完 d 之后加 printf(" %s: %.2f\n", db->rows[i].name, d);。观察输出:距离小的记录排在前面,距离大的可能被后面的更小距离替换掉。

Q5:两种做法结果对比。 用同一组数据分别调用 knn_sortknn_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 聚类 要解决的场景。