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

基于回归的图像类相似度占比计算及代码异常排查

问题

我正在做一个针对雏菊(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. 如何用回归算法实现图像与各类别的相似度百分比输出?
  2. 当前代码导致程序冻结、图像无法显示的问题出在哪里?

解决方案

一、用回归算法实现相似度百分比输出

针对二分类场景,可采用一对多回归的思路实现:

  • 把原标签转换为两个回归目标:比如雏菊类对应目标[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)才有,这是你之前代码异常的核心原因之一。

二、解决程序冻结与图像显示问题

当前代码的主要问题及修复方案:

  1. 错误调用predict_proba:回归器没有该方法,调用会直接抛出AttributeError,导致后续代码中断。需替换为回归器的predict方法,按上面的逻辑生成相似度。
  2. 图像更新逻辑缺失:每次显示新图像时未清除之前的绘图内容,会导致图像叠加、界面卡顿。需在display_image方法开头添加self.ax.clear()。
  3. 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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.03 23:40:57