Sklearn NearestNeighbors使用mahalanobis距离查询数组切片报错如何解决
报错原因
你虽然没有显式传入Y参数,但当kneighbors()方法传入的查询样本集和fit()阶段使用的训练集不是完全相同的数组时,scikit-learn内部会自动将训练集标记为X、查询集标记为Y,进入双样本距离计算逻辑,因此触发马氏距离的参数校验规则。
而你传入整个X作为查询集时,scikit-learn识别到查询集和训练集完全一致,会走单样本内部优化逻辑,不需要额外的逆矩阵参数,因此不会报错。
解决方法
马氏距离计算本质上需要用到协方差矩阵的逆矩阵,你直接在metric_params中传入协方差矩阵的逆(对应参数名VI)即可兼容单查询/批量查询的所有场景,修改后的代码如下:
import numpy as np from sklearn.datasets import make_classification from sklearn.neighbors import NearestNeighbors X, y = make_classification() # 计算协方差矩阵的逆作为VI参数传入 cov_matrix = np.cov(X, rowvar=False) nn = NearestNeighbors(algorithm='brute', metric='mahalanobis', metric_params={'VI': np.linalg.inv(cov_matrix)}) # 此时查询单个样本也可以正常运行 res = nn.fit(X).kneighbors(X[:1, :])
如果你的特征存在多重共线性导致协方差矩阵不可逆,可在求逆前给协方差矩阵对角元素加极小的正则值规避报错:
cov_matrix = np.cov(X, rowvar=False) + 1e-6 * np.eye(X.shape[1])
内容的提问来源于stack exchange,提问作者Yandle
相关产品推荐
相关产品推荐

