如何修正代码以生成索引对应class_weight值的正确字典?
问题描述
我有如下示例数据集:
cat count class_weight cars 17824 0.000404 bus 124784 0.000553 planes 111271 0.000620
我希望输出一个包含索引和class_weight值的dictionary,格式如下:
{0: 0.000404, 1: 0.000553, 2: 0.000620}
我当前的实现代码为:
class_weight = {} for index, label in enumerate(categories): class_weight[index] = df[df['cat'] == categories]['class_weight'].values[0]
其中categories是存储cat列值的numpy.ndarray。但输出结果错误,所有值都返回第一个元素:
{0: 0.000404, 1: 0.000404, 2: 0.000404}
请问如何修正这段代码?
修正方案
1. 修复循环逻辑错误
你的代码核心问题是筛选条件写错了:用categories(整个数组)和df['cat']比较,会导致每次筛选都匹配所有行,取values[0]自然只会得到第一个元素。把筛选条件改成当前循环的label即可:
class_weight = {} for index, label in enumerate(categories): class_weight[index] = df[df['cat'] == label]['class_weight'].values[0]
2. 更高效的Pandas原生实现
如果不需要手动循环,直接用Pandas内置方法可以一步生成目标字典:
- 方法一:重置索引后转字典
class_weight = df.reset_index(drop=True)['class_weight'].to_dict()
reset_index(drop=True)会将原数据行索引重置为0开始的连续整数,之后提取class_weight列直接转成字典即可。
- 方法二:直接枚举列值
如果categories的顺序和df中cat列的顺序完全一致,还可以更简洁:
class_weight = dict(enumerate(df['class_weight'].values))
内容的提问来源于stack exchange,提问作者Katty_one
相关产品推荐
相关产品推荐

