图像生成器类别频率可视化时遇numpy.float64不可迭代TypeError
解决类别频率可视化时的TypeError问题
问题场景
构建图像生成器时使用了以下代码:
train_data_dir= "/Users/awabe/Desktop/Project/PapilaDB/ExpertsSegmentations/ImagesWithContours train" train_datagen = ImageDataGenerator(rescale=1./255, shear_range=0.2, zoom_range=0.2, horizontal_flip=True, validation_split=0.2) # set validation split train_generator = train_datagen.flow_from_directory( train_data_dir, target_size=(2576, 2000), batch_size=24, class_mode='binary', subset='training') # set as training data validation_generator = train_datagen.flow_from_directory( train_data_dir, # same directory as training data target_size=(2576, 2000), batch_size=24, class_mode='binary', subset='validation') # set as validation data
随后绘制类别频率柱状图:
plt.xticks(rotation=90) plt.bar(x=labels, height=np.mean(train_generator.labels, axis=0)) plt.title("Frequency of Each Class") plt.show()
得到每个类别的频率为1.5后,使用compute_class_freqs获取正负类别频率:
freq_pos, freq_neg = compute_class_freqs(train_generator.labels) freq_pos
但执行以下可视化代码时触发错误:
data = pd.DataFrame({"Class": labels, "Label": "Positive", "Value": freq_pos}) data = data.append([{"Class": labels[l], "Label": "Negative", "Value": v} for l,v in enumerate(freq_neg)], ignore_index=True) plt.xticks(rotation=90) f = sns.barplot(x="Class", y="Value", hue="Label" ,data=data)
错误信息:
--------------------------------------------------------------------------- TypeError Traceback (most recent call last) Cell In[173], line 2 1 data = pd.DataFrame({"Class": labels, "Label": "Positive", "Value": freq_pos}) ----> 2 data = data.append([{"Class": labels[l], "Label": "Negative", "Value": v} for l,v in enumerate(freq_neg)], ignore_index=True) 3 plt.xticks(rotation=90) 4 f = sns.barplot(x="Class", y="Value", hue="Label" ,data=data) TypeError: 'numpy.float64' object is not iterable
错误原因
报错核心是freq_neg是单个numpy.float64数值,而enumerate(freq_neg)会尝试迭代这个浮点数,浮点数无法被迭代,因此触发TypeError。
因为你使用的是二分类(class_mode='binary'),compute_class_freqs返回的freq_pos和freq_neg都是单个数值,分别对应正类和负类的频率,并非可迭代的数组或列表。
解决方案
无需使用列表推导式循环,直接为每个类别构造正负频率的行数据即可。以下是修正后的代码:
方式一:直接构造完整的DataFrame
# 假设labels是类名列表,比如['Class0', 'Class1'],对应负类和正类 data = pd.DataFrame([ # 负类的正负频率 {"Class": labels[0], "Label": "Positive", "Value": freq_neg}, {"Class": labels[0], "Label": "Negative", "Value": 1 - freq_neg}, # 正类的正负频率 {"Class": labels[1], "Label": "Positive", "Value": freq_pos}, {"Class": labels[1], "Label": "Negative", "Value": 1 - freq_pos} ]) plt.xticks(rotation=90) f = sns.barplot(x="Class", y="Value", hue="Label", data=data) plt.show()
方式二:分步添加数据
如果需要保留原代码的分步逻辑,可以修改为:
# 初始化正类数据 data = pd.DataFrame({ "Class": labels, "Label": ["Positive"] * len(labels), "Value": [freq_neg, freq_pos] # 对应labels中两个类的正频率 }) # 添加负类数据 neg_data = pd.DataFrame({ "Class": labels, "Label": ["Negative"] * len(labels), "Value": [1 - freq_neg, 1 - freq_pos] # 对应两个类的负频率 }) data = pd.concat([data, neg_data], ignore_index=True) plt.xticks(rotation=90) f = sns.barplot(x="Class", y="Value", hue="Label", data=data) plt.show()
注意:如果compute_class_freqs的返回值定义不同(比如返回的是每个样本的正负频率数组),需要根据实际逻辑调整,但针对二分类场景,上述代码可以解决当前报错。
内容的提问来源于stack exchange,提问作者Awab Elkhair
相关产品推荐
相关产品推荐

