关于numpy ndarray索引及鸢尾花数据集预测编码转标签的技术疑问
解决鸢尾花数据集预测结果与标签映射的问题
嘿,我完全懂你在鸢尾花实践里遇到的困惑!让我一步步给你拆解清楚:
核心需求:把数字预测结果转成品种名称
你手里的两个ndarray刚好是完美配对的:
iris.target_names是按顺序存储品种名称的一维数组,索引0对应'setosa',1对应'versicolor',2对应'virginica'clf.predict(test[features])是模型输出的数字编码数组,每个数字对应上述的索引
要把数字转成对应的名称,直接用预测数组作为索引去访问iris.target_names就行,代码超简单:
predicted_names = iris.target_names[clf.predict(test[features])]
举个小例子验证:如果你的预测结果是array([0, 0, 1, 2]),运行上面的代码后,predicted_names就会输出array(['setosa', 'setosa', 'versicolor', 'virginica'], dtype='<U10'),完全符合预期!
关于numpy ndarray索引的补充说明
这种用一个ndarray去索引另一个ndarray的方式叫整数数组索引,是numpy里非常实用的功能:
- 索引数组的形状可以和原数组不同,但每个元素必须是原数组的有效索引(不能超出原数组的范围,比如这里原数组长度是3,索引不能是3或负数,除非你用负数索引取倒数)
- 返回的结果数组会和索引数组的形状完全一致,每个位置的元素就是原数组对应索引的取值
再举个更通用的例子帮你理解:
import numpy as np # 原数组:存储类别名称 categories = np.array(['cat', 'dog', 'bird']) # 索引数组:存储类别编码 labels = np.array([1, 0, 2, 1, 1]) # 用索引数组获取对应名称 result = categories[labels] print(result) # 输出:array(['dog', 'cat', 'bird', 'dog', 'dog'], dtype='<U4')
这样是不是就清晰多啦?
内容的提问来源于stack exchange,提问作者ash bounty
相关产品推荐
相关产品推荐

