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

如何使用NumPy向量化函数替代循环计算鸢尾花分类准确率?

用NumPy向量化替代K近邻分类中的循环

原问题代码

用户实现了一个基于K近邻(K=1)的分类准确率计算程序,但嵌套循环效率较低,希望用NumPy向量化操作替代所有循环:

#command to import NumPy package
import numpy as np


iris_train=np.genfromtxt("iris-train-data.csv",delimiter=',',usecols=(0,1,2,3),dtype=float)
iris_test=np.genfromtxt("iris-test-data.csv",delimiter=',',usecols=(0,1,2,3),dtype=float)


train_cat=np.genfromtxt("iris-training-data.csv",delimiter=',',usecols=(4),dtype=str)
test_cat=np.genfromtxt("iris-testing-data.csv",delimiter=',',usecols=(4),dtype=str)


correct = 0

for i in range(len(iris_test)):
    n = 0
    old_distance = float('inf')
    
    
    while n < len(iris_train):
        #finding the difference between test and train point
        iris_diff = (abs(iris_test[i] - iris_train[n])**2)
        #summing up the calculated differences
        iris_sum = sum(iris_diff)
        new_distance = float(np.sqrt(iris_sum))
        
        #if statement to update distance
        if new_distance < old_distance:
            index = n
            old_distance = new_distance
        n += 1
        
    
    print(i + 1, test_cat[i], train_cat[index])
    if test_cat[i] == train_cat[index]:
        correct += 1
        

accuracy = ((correct)/float((len(iris_test)))*100)
print(f"Accuracy:{accuracy: .2f}%")

向量化优化后的代码

利用NumPy的广播机制和矩阵运算,可以完全去除循环,大幅提升效率:

import numpy as np

# 加载数据(修正原代码中训练标签的文件名,与训练数据文件名保持一致)
iris_train = np.genfromtxt("iris-train-data.csv", delimiter=',', usecols=(0,1,2,3), dtype=float)
iris_test = np.genfromtxt("iris-test-data.csv", delimiter=',', usecols=(0,1,2,3), dtype=float)
train_cat = np.genfromtxt("iris-train-data.csv", delimiter=',', usecols=(4), dtype=str)
test_cat = np.genfromtxt("iris-test-data.csv", delimiter=',', usecols=(4), dtype=str)

# 计算所有测试点与训练点的欧氏距离平方(无需开根号,不影响最小值判断)
diff = iris_test[:, np.newaxis] - iris_train
squared_diff = diff ** 2
distance_squared = squared_diff.sum(axis=2)

# 一次性找出每个测试点对应的最近训练点索引
min_indices = np.argmin(distance_squared, axis=1)

# 获取每个测试点的预测类别
predicted_cats = train_cat[min_indices]

# 打印每个测试点的结果
for idx, (true_cat, pred_cat) in enumerate(zip(test_cat, predicted_cats), start=1):
    print(idx, true_cat, pred_cat)

# 计算准确率
correct = np.sum(test_cat == predicted_cats)
accuracy = (correct / len(iris_test)) * 100
print(f"Accuracy:{accuracy: .2f}%")

关键优化点说明

  • 广播机制:通过iris_test[:, np.newaxis]将测试数据形状从(N_test,4)转为(N_test,1,4),与训练数据(N_train,4)运算时自动广播为(N_test,N_train,4),一次性计算所有测试点与训练点的差值。
  • 向量化求和:用sum(axis=2)直接对每个测试点-训练点对的特征维度求和,替代循环中的手动累加。
  • argmin批量找最近邻:np.argmin(distance_squared, axis=1)一次性找出每个测试点对应的最近训练点索引,替代嵌套循环的距离比较逻辑。
  • 省略开根号:欧氏距离的平方和与原距离的大小关系一致,直接用平方和找最小值可以减少计算量,不影响结果。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.02 22:20:41