如何修正k近邻预测函数代码以获取长度为k的索引列表
k近邻预测函数的索引获取逻辑修正
原代码问题分析
你原来的问题行存在两个核心错误:
- 每次循环都重复调用
find_index,返回的永远是x_new所在区间的同一个左边界索引,最终得到的nk列表全是重复值,完全不符合k近邻取k个最近点的要求 - 列表推导式的循环变量用了和入参相同的
k,会意外覆盖入参的k值
正确修正方案
利用x坐标已经升序排列的特性,以find_index返回的基准索引为中心,往左右两侧扩展,每次选择离x_new更近的索引加入结果列表,直到凑够k个索引即可。
修正后的完整knn_predict函数代码如下:
def knn_predict(data, x_new, k): """ (tuple, number, int) -> number data is a tuple. data[0] are the x coordinates and data[1] are the y coordinates. k is a positive nearest neighbor parameter. Returns k-nearest neighbor estimate using nearest neighbor parameter k at x_new. Assumes i) there are no duplicated values in data[0], ii) data[0] is sorted in ascending order, and iii) x_new falls between min(x) and max(x). >>> knn_predict(([0, 5, 10, 15], [1, 7, -5, 11]), 2, 2) 4.0 >>> knn_predict(([0, 5, 10, 15], [1, 7, -5, 11]), 2, 3) 1.0 >>> knn_predict(([0, 5, 10, 15], [1, 7, -5, 11]), 8, 2) 1.0 >>> knn_predict(([0, 5, 10, 15], [1, 7, -5, 11]), 8, 3) 4.333333333333333 """ x_list = data[0] y_list = data[1] # 获取x_new所在区间的左边界基准索引 base_idx = find_index(x_list, x_new) nk = [] # 初始化左右指针 left_ptr = base_idx right_ptr = base_idx + 1 for _ in range(k): # 优先选择离x_new更近的索引,边界越界时直接选另一侧 if left_ptr >= 0 and (right_ptr >= len(x_list) or abs(x_list[left_ptr] - x_new) <= abs(x_list[right_ptr] - x_new)): nk.append(left_ptr) left_ptr -= 1 else: nk.append(right_ptr) right_ptr += 1 # 计算k个近邻的y值均值 yvals = [y_list[val] for val in nk] ynew = sum(yvals) / k return ynew
验证说明
上述代码完全匹配你给出的所有测试用例结果,逻辑是利用有序数组的特性,无需排序所有距离即可高效拿到k个最近点的索引,时间复杂度为O(k),优于全量计算距离再排序的O(nlogn)方案。
内容的提问来源于stack exchange,提问作者ashnotallyson
相关产品推荐
相关产品推荐

