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

不使用sklearn手动实现欧氏距离KNN 调整k值输出多组准确率

K近邻多k值准确率计算修改方案

原有代码问题说明

  • classtrain取值错误:原代码错误从测试集取训练样本的标签,应该改为从训练集对应行取
  • 邻居存储逻辑错误:原逻辑仅保留1个最近邻居,无法支持多k值的统计需求
  • 变量名拼写错误:训练集花萼宽度变量名拼写错误,会导致计算异常

修改后完整代码

import math

# 数据集打乱拆分逻辑保留
shuffle_df = dataset.sample(frac=1)
train_size = int(0.75 * len(dataset))
train_set = shuffle_df[:train_size]
test_set = shuffle_df[train_size:]

# 要测试的k值列表
k_list = [1, 3, 5]
accuracy_result = {}

for k in k_list:
    correct_count = 0
    # 遍历所有测试样本
    for w in range(len(test_set)):
        # 取测试样本的特征和真实标签
        sepallengthtest = test_set.iloc[w,0]
        sepalwidthtest = test_set.iloc[w,1]
        petallenghttest = test_set.iloc[w,2]
        petalwidthtest = test_set.iloc[w,3]
        true_label = test_set.iloc[w,4]
        
        # 存储当前测试样本和所有训练样本的距离+训练样本标签
        distance_list = []
        for m in range(len(train_set)):
            # 取训练样本的特征和标签
            sepallengthtrain = train_set.iloc[m,0]
            sepalwidthtrain = train_set.iloc[m,1]
            petallenghttrain = train_set.iloc[m,2]
            petalwidthtrain = train_set.iloc[m,3]
            train_label = train_set.iloc[m,4]
            
            # 计算欧式距离
            distance = math.sqrt(
                (sepallengthtest - sepallengthtrain)**2 + 
                (sepalwidthtest - sepalwidthtrain)**2 + 
                (petallenghttest - petallenghttrain)**2 + 
                (petalwidthtest - petalwidthtrain)**2
            )
            distance_list.append((distance, train_label))
        
        # 按距离升序排序,取前k个最近邻居
        distance_list.sort(key=lambda x:x[0])
        top_k_neighbors = distance_list[:k]
        
        # 统计邻居的标签投票数
        vote_dict = {}
        for item in top_k_neighbors:
            label = item[1]
            vote_dict[label] = vote_dict.get(label, 0) + 1
        
        # 取得票最高的标签作为预测结果
        predict_label = max(vote_dict.items(), key=lambda x:x[1])[0]
        if predict_label == true_label:
            correct_count += 1
    
    # 计算当前k的准确率
    accuracy = correct_count / len(test_set)
    accuracy_result[k] = accuracy

# 按要求格式输出
print(f"{'Neighbors:':<15}{k_list[0]:^8}{k_list[1]:^8}{k_list[2]:^8}")
print(f"{'Success Rate:':<15}{accuracy_result[k_list[0]]:>7.1%}{accuracy_result[k_list[1]]:>8.1%}{accuracy_result[k_list[2]]:>8.1%}")

代码说明

  • 通用化读取拆分后的数据集长度,避免硬编码数值导致数据集变动时报错
  • 新增标签投票逻辑,通过统计前k个邻居的标签出现频次确定预测结果
  • 输出格式做了对齐处理,和需求的期望示例格式完全匹配

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.27 19:45:03