基于回归的图像类相似度占比计算及代码异常排查
问题
我正在做一个针对雏菊(daisies)和向日葵(sunflowers)两类花卉数据集的项目,要求使用分类与回归算法。目前已经用随机森林分类实现了测试图像的类别判定,但作业要求必须用回归算法输出图像与每类花卉的相似度百分比(比如某雏菊图像显示87%雏菊、13%向日葵)——虽然知道回归不是这个场景的最优方案,但必须按要求执行。
我尝试用Random Forest Regressor和scikit-learn的predict_proba方法,结果程序直接冻结,而且添加回归相关代码后,原本正常的图像显示功能也失效了。相关代码片段如下:
# Convert lists to numpy arrays test_images = np.array(test_images) test_labels = np.array(test_labels) test_probabilities = classifier.predict_proba(test_images) # Create an interactive image viewer for the test set class ImageViewer: def __init__(self, images, true_labels, predicted_probabilities): self.images = images self.true_labels = true_labels self.predicted_probabilities = predicted_probabilities self.index = 0 self.fig, self.ax = plt.subplots() self.display_image() self.next_button = Button(plt.axes([0.7, 0.02, 0.1, 0.05]), 'Next') self.next_button.on_clicked(self.next_image) plt.show() def display_image(self): img = self.images[self.index].reshape(256, 256, 3) true_class = self.true_labels[self.index] predicted_probs = self.predicted_probabilities[self.index] self.ax.imshow(cv2.cvtColor(img, cv2.COLOR_BGR2RGB)) self.ax.set_title(f"True: {true_class}\nPredicted Probabilities: {predicted_probs}") self.ax.axis('off') def next_image(self, event): self.index = (self.index + 1) % len(self.images) self.display_image() plt.draw() # Initialize the image viewer for the test set test_image_viewer = ImageViewer(test_images, test_labels, test_probabilities)
想请教两个问题:
- 如何用回归算法实现图像与各类别的相似度百分比输出?
- 当前代码导致程序冻结、图像无法显示的问题出在哪里?
解决方案
一、用回归算法实现相似度百分比输出
针对二分类场景,可采用一对多回归的思路实现:
- 把原标签转换为两个回归目标:比如雏菊类对应目标
[1, 0],向日葵类对应目标[0, 1] - 训练两个
RandomForestRegressor模型:一个预测样本属于雏菊的“相似度”,另一个预测属于向日葵的“相似度” - 预测后对两个模型的输出做归一化处理,转换成百分比(若输出和不为1,就除以两者之和得到占比)
示例代码片段:
from sklearn.ensemble import RandomForestRegressor import numpy as np # 假设训练集图像已扁平化(train_images_flat),训练标签为0=雏菊、1=向日葵 # 构建回归目标 train_target_daisy = (train_labels == 0).astype(float) train_target_sunflower = (train_labels == 1).astype(float) # 训练两个回归器 reg_daisy = RandomForestRegressor(n_estimators=100, random_state=42) reg_daisy.fit(train_images_flat, train_target_daisy) reg_sunflower = RandomForestRegressor(n_estimators=100, random_state=42) reg_sunflower.fit(train_images_flat, train_target_sunflower) # 预测测试集(先扁平化测试图像) test_images_flat = test_images.reshape(test_images.shape[0], -1) pred_daisy = reg_daisy.predict(test_images_flat) pred_sunflower = reg_sunflower.predict(test_images_flat) # 归一化得到百分比 total = pred_daisy + pred_sunflower prob_daisy = (pred_daisy / total) * 100 prob_sunflower = (pred_sunflower / total) * 100 # 组合成相似度结果 test_probabilities = np.column_stack((prob_daisy, prob_sunflower))
注意:RandomForestRegressor没有predict_proba方法,只有分类器(RandomForestClassifier)才有,这是你之前代码异常的核心原因之一。
二、解决程序冻结与图像显示问题
当前代码的主要问题及修复方案:
- 错误调用
predict_proba:回归器没有该方法,调用会直接抛出AttributeError,导致后续代码中断。需替换为回归器的predict方法,按上面的逻辑生成相似度。 - 图像更新逻辑缺失:每次显示新图像时未清除之前的绘图内容,会导致图像叠加、界面卡顿。需在
display_image方法开头添加self.ax.clear()。 - Matplotlib交互模式冲突:默认
plt.show()会阻塞主线程,结合按钮回调易引发冻结。需开启交互模式plt.ion(),并调整plt.show()的调用参数。
修改后的ImageViewer类:
class ImageViewer: def __init__(self, images, true_labels, predicted_probabilities): self.images = images self.true_labels = true_labels self.predicted_probabilities = predicted_probabilities self.index = 0 plt.ion() # 开启交互模式 self.fig, self.ax = plt.subplots() self.display_image() self.next_button = Button(plt.axes([0.7, 0.02, 0.1, 0.05]), 'Next') self.next_button.on_clicked(self.next_image) plt.show(block=True) # 阻塞保持窗口打开 def display_image(self): self.ax.clear() # 清除之前的绘图内容 img = self.images[self.index].reshape(256, 256, 3) true_class = self.true_labels[self.index] predicted_probs = self.predicted_probabilities[self.index] # 格式化百分比显示 prob_text = f"Daisy: {predicted_probs[0]:.1f}%, Sunflower: {predicted_probs[1]:.1f}%" self.ax.imshow(cv2.cvtColor(img, cv2.COLOR_BGR2RGB)) self.ax.set_title(f"True: {true_class}\n{prob_text}") self.ax.axis('off') self.fig.canvas.draw() # 强制更新画布 def next_image(self, event): self.index = (self.index + 1) % len(self.images) self.display_image()
另外,必须确保测试图像已扁平化后再输入回归模型,否则三维图像(256,256,3)会导致回归器计算量暴增,引发程序冻结。
内容的提问来源于stack exchange,提问作者andreeam
相关产品推荐
相关产品推荐

