2.2 现场五:手写KNN与距离度量


2.2 现场五:手写 KNN 与距离度量

本节摘要:手写 K 近邻分类器:广播一次算出全部距离、argsort 取近邻、多数投票出标签。追问链压向距离公式的展开形式、高维失效与 KD 树加速。这道题考查"推断也可以没有训练"的范式,以及 numpy 思维——循环写在数组运算里,不写在 for 里。

上一节练的是训练循环,这一节反过来:KNN 没有任何可训练参数,所有计算都推迟到预测时刻。两种范式对照着写一遍,比背十遍定义更能理解"模型"这个词的弹性。

面试官提问

"实现 KNN 分类:给我训练集和查询点,返回预测标签。不许用 sklearn。写完解释你的距离计算为什么没有双重循环。"

最后一问是本题的题眼。双重循环版本的 KNN 人人会写,面试官要看的是你能不能用广播把循环压进 C 层。

现场推演

候选人先口述:"对每个查询点,算它到全部训练点的欧氏距离,取最近的 k 个,看哪类标签最多。"然后主动升级:"如果查询点很多,逐个算就慢了,我用广播一次算出整个距离矩阵。"落笔:

import numpy as np from collections import Counter train = np.array([[1.0, 1.0], [1.2, 0.8], [5.0, 5.0], [5.2, 4.8], [4.9, 5.1]]) labels = np.array(["A", "A", "B", "B", "B"]) query = np.array([[1.1, 0.9], [5.1, 5.0]]) def knn_predict(train, labels, query, k=3): # 广播:query(2,1,2) - train(5,2) -> (2,5,2),平方沿最后轴求和 d = np.sqrt(((query[:, None, :] - train[None, :, :]) ** 2).sum(axis=2)) preds = [] for i in range(len(query)): idx = np.argsort(d[i])[:k] # 最近的 k 个下标 top = Counter(labels[idx]).most_common(1)[0][0] preds.append(top) return np.array(preds) print(knn_predict(train, labels, query))
['A' 'B']

两个查询点分别落在两簇中心,预测 A 与 B,直觉一致。候选人指着 query[:, None, :] - train[None, :, :] 解释:"中间插入的新轴让两个矩阵广播成逐对之差,三行代码替代了双重循环,计算落在 numpy 的 C 实现里。"

追问链

第一问:距离公式能展开吗? 欧氏距离的平方等于 x 平方加 y 平方减二倍内积。展开后查询点之间的计算全部变成矩阵乘法,这正是大规模近邻检索的标准加速写法。

# 展开式:d 等于 q范数平方 + t范数平方 - 2倍的点积 q_sq = (query ** 2).sum(axis=1, keepdims=True) # (2,1) t_sq = (train ** 2).sum(axis=1) # (5,) cross = query @ train.T # (2,5) d2 = q_sq + t_sq - 2 * cross print(d2.round(3))
[[ 0.02 0.09 32.43 32.9 32.5 ] [32.42 31.69 0.02 0.05 0.02]]

与广播版逐元素一致(sqrt 前的平方距离),且矩阵乘法可以调底层 BLAS。候选人补了一句取舍:"展开式有数值技巧问题——浮点误差可能让 d2 出现极小的负数,工程上要 clip 到零再开根。"

第二问:k 怎么选?票数相同呢? 候选人答:"k 通常用交叉验证扫一遍;票数平票时常见约定是取距离更近的那类,或者 k 取奇数在二分类里直接避免平票。"现场改 k=1 验证退化行为:

print(knn_predict(train, labels, query, k=1))
['A' 'B']

k 为一时 KNN 完全依赖最近一个点,对噪声敏感——这句话是通往"偏差方差"追问的门票。

第三问:维度很高时 KNN 还好用吗? 这是本题最深的一层。候选人答:"高维空间里所有点对之间的距离趋于接近,'最近'失去区分度,这就是维度灾难。缓解办法:先降维(PCA)再找近邻,或者用近似最近邻方法牺牲一点精度换数量级的加速,比如基于哈希分桶或图的检索结构。"面试官追问复杂度:"暴力是每次查询乘以训练集规模;树结构在低维能把单次查询压到对数级,但维度一高树会退化回线扫。"

第四问:换个距离度量呢? 面试官把查询点指回那两个簇:"用余弦相似度找近邻,结果还一样吗?"候选人算给他看:

q = np.array([1.1, 0.9]) t1, t2 = np.array([1.0, 1.0]), np.array([5.0, 5.0]) def cos(u, v): return float(u @ v / (np.linalg.norm(u) * np.linalg.norm(v))) print('cos(q, 近簇):', round(cos(q, t1), 3), ' cos(q, 远簇):', round(cos(q, t2), 3)) print('L1 到近簇:', round(float(np.abs(q - t1).sum()), 1), ' L1 到远簇:', round(float(np.abs(q - t2).sum()), 1))
cos(q, 近簇): 0.995 cos(q, 远簇): 0.995 L1 到近簇: 0.2 L1 到远簇: 8.0

余弦把两个方向的夹角都算成零点九九五——模长被扔掉了,远近无从区分;欧氏与曼哈顿都毫不含糊地选中近簇。他把结论收成一句:"文本检索常用余弦,因为文档只关心词的配比方向;几何坐标用欧氏,因为模长本身就是信息。距离度量是先验的编码,不是随手挑的函数。"面试官在这句上记了一笔——这层理解已经越过"会写 KNN",摸到"懂表示"了。

失误复盘

高频翻车点:广播维度拼错,query 与 train 的形状对调,距离矩阵行列含义颠倒,结果"看起来也能跑"但全是错的——argsort 沿错误轴取近邻不会报错,只会安静地错;开根号忘了或平方忘了,排序结果不变(单调变换不改变近邻序),但打印"距离值"时露馅;投票用 set 去重而不是 Counter 计数,平票行为随机。还有一个白板细节:训练集只有五个点却把 k 设成五以上,取到的近邻数不足 k,Counter 行为与预期不符。

主线候选人这场发挥最好的一次是展开式追问——他不仅写了展开式,还主动提到数值负数问题。面试官事后说,这一句"工程上要 clip"值半轮好评。

关键直觉:广播与矩阵展开把 Python 循环换成数组运算,这一招从 KNN 一路用到第 3 章卷积的 im2col——数值计算的加速思路是同一棵树上长出来的。


作者与出处
原作者: 灏天文库
来源:灏天文库
整理: 灏天文库整理
由灏天文库平台收录,内容或由平台用户上传,仅供学习交流
发布者: 作者: 灏天文库 转发
评论区 (0)
U