You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

sklearn中实现带权重向量的自定义距离函数及报错解决

实现Scikit-Learn带权自定义距离

需求为实现支持样本权重参数的自定义汉明距离,距离计算规则如下:
加权距离计算规则
原有实现运行抛出错误:TypeError: Custom distance function must accept two vectors and return a float.

问题根因

原有代码触发报错有三个核心问题:

  • 方法引用错误:类内定义的距离方法是weighted_hamming,传入DistanceMetric时错写成了不存在的hamming方法
  • 变量作用域错误:距离计算逻辑里写的w是全局变量,没有调用实例绑定的self.w,且直接和sum(a!=b)相乘返回的是数组,不符合函数必须返回单个浮点数的要求
  • 权重逻辑不匹配:传入的权重形状是(4,1),且没有做样本和权重的对应映射,计算两个样本距离时无法正确取到对应的权重值。

Scikit-Learn对pyfunc自定义距离的强制要求是:传入的可调用对象只能接收两个一维数组参数(待计算距离的两个样本),且必须返回单个float类型的标量值,额外的权重参数需要通过闭包/类实例绑定的方式传入,不能直接作为第三个参数塞给距离函数。

正确实现代码

from sklearn.neighbors import DistanceMetric
import numpy as np

class WeightedHamming:
    def __init__(self, sample_weights):
        # 入参sample_weights为一维数组,长度等于样本总数,对应每个样本的权重
        self.w = sample_weights
        self.feature_cache = None

    def _calc_dist(self, a, b):
        # 匹配当前两个样本对应的索引,取权重
        idx_a = np.where((self.feature_cache == a).all(axis=1))[0][0]
        idx_b = np.where((self.feature_cache == b).all(axis=1))[0][0]
        # 统计不相等的特征数量
        mismatch = np.sum(a != b)
        # 按加权规则计算距离,最终强转为float类型返回
        # 注:如果你的权重是特征维度而非样本维度,不需要做索引匹配,直接给每个不匹配特征乘对应权重求和即可
        dist = mismatch * (self.w[idx_a] + self.w[idx_b]) / 2
        return float(dist)

    def get_metric(self, X):
        self.feature_cache = X
        return DistanceMetric.get_metric(metric='pyfunc', func=self._calc_dist)

# 测试用例
if __name__ == "__main__":
    # 汉明距离适用于离散特征,这里用0/1离散值生成测试数据
    X = np.random.randint(0, 2, size=(4, 3))
    # 权重调整为一维数组,长度和样本数一致
    w = np.random.random(4)
    cal = WeightedHamming(w)
    metric = cal.get_metric(X)
    print(metric.pairwise(X))

使用注意事项

  • 上述实现适配每个样本对应一个权重的需求,如果是给K近邻分类/回归器用,预测阶段的新样本需要提前补充对应权重到缓存中,否则会出现索引匹配失败的问题
  • 距离函数的返回值必须强转为float类型,禁止返回数组、numpy标量等其他类型,否则会触发类型报错
  • 如果权重是每个特征对应一个权重(而非每个样本对应权重),可以简化逻辑:初始化时传入特征权重数组,计算时给每个不相等的特征乘对应权重再求和即可,不需要做样本索引匹配。

内容的提问来源于stack exchange,提问作者LeparaLaMapara

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.26 15:45:33