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

K近邻(KNN)实现对鸢尾花数据集预测始终返回相同类别求助

KNN代码错误定位及修复

以下是代码中存在的4个核心问题:

  • 数组越界错误
    你定义训练集数组为double train_data[15][3],仅支持索引0、1、2三个维度,但后续将距离存入索引3的位置,同时排序时遍历到索引3,直接触发内存越界,数据读写结果不可控。需要修改数组定义为double train_data[15][4]。
  • 排序逻辑反向
    你需要筛选距离待预测点最近的K个样本,应该按距离升序排列(近的在前),但当前代码判断train_data[i][3] < train_data[j][3]时交换样本,实际是按距离降序排列,取到的前K个都是距离最远的样本,计算逻辑完全错误。需要将判断条件改为train_data[i][3] > train_data[j][3]。
  • 最大值索引函数逻辑错误
    findMaxIndex函数中已经计算出了计数最多的类别索引maxIndex,但函数末尾写死返回0,导致不管统计结果如何,最终返回的索引永远是0,这是程序始终输出类别1的直接原因。需要将返回值改为return maxIndex。
  • 类别索引偏移错误
    你定义的类别标签本身就是0/1/2,和初始化的classes数组索引刚好对应,但统计计数时写了classes[(int)train_data[i][2]-1],当样本类别为0时,索引会变成-1触发数组越界。需要去掉-1的偏移,后续预测标签也不用额外加1。

修复后的完整代码如下:

#include <iostream>
#include <math.h>
#include <string>

//Setosa = 0, Virginica = 1, Verscicolor = 2
//[0] and [1] = data point, [2] = class, [3] = distance
double train_data[15][4] = {
{5.3,3.7,0},{5.1,3.8,0},{7.2,3.0,1},
{5.4,3.4,0},{5.1,3.3,0},{5.4,3.9,0},
{7.4,2.8,1},{6.1,2.8,2},{7.3,2.9,1},
{6.0,2.7,2},{5.8,2.8,1},{6.3,2.3,2},
{5.1,2.5,2},{6.3,2.5,2},{5.5,2.4,2}
};

double Distance(double attr1, double attr2, double sAttr1, double sAttr2)
{
    return sqrt(pow(attr1-sAttr1, 2.0)+pow(attr2-sAttr2, 2.0));
}

int findMaxIndex(float *classes)
{
    int maxIndex = 0;
    for (int i = 0; i < 3; i++){
        if (classes[i] > classes[maxIndex])
        maxIndex = i;
    }
    return maxIndex;
}

int main(){
    for(int i = 0; i < 15; i++){
        train_data[i][3] = Distance(train_data[i][0],train_data[i][1],5.2,3.1);
    }

    for(int i = 0; i < 15; i++){
        for (int j = i+1; j < 15; j++){
            if (train_data[i][3] > train_data[j][3]){
                //swap
                for(int k = 0; k < 4; k++){
                    double temp = train_data[i][k];
                    train_data[i][k] = train_data[j][k];
                    train_data[j][k] = temp;
                }
            }
        }
    }   

    //Based on a value for k determine the class
    int K = 5;
    float *classes = new float[3];
    for (int i =0; i < 3; i++){
        classes[i] = 0;
    }
    for (int i = 0 ; i < K; i++)
    {
        classes[(int)train_data[i][2]]++;
    }
    
    int predictedLabel = findMaxIndex(classes);
    std::cout << "Predicted class for point {5.2,3.1} is: " << predictedLabel << std::endl;
    return 0;
}

修复后运行输出为0,也就是对应Setosa类别,符合K近邻计算的预期结果。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.28 11:36:03