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

使用yellowbrick绘制回归预测误差图时OLS、Keras模型报错如何解决?

报错原因

Yellowbrick的回归类可视化器完全基于scikit-learn的API规范设计,会调用sklearn内置的is_regressor()函数校验传入模型的类型,只有符合以下两个条件才会被判定为合法回归器:

  • 模型带有_estimator_type = 'regressor'的属性标记
  • 模型遵循sklearn估计器的统一接口(有fit、predict等标准方法,调用逻辑与sklearn一致)

你用到的三个模型中,RandomForestRegressor是sklearn原生实现,天然符合上述要求,所以可以正常出图;另外两个模型不满足规范,因此触发报错:

  1. statsmodels的sm.OLS()实现的线性回归,没有遵循sklearn的接口设计,既没有内置_estimator_type标记,拟合、预测的调用逻辑也和sklearn不匹配,无法通过校验。
  2. 基于TensorFlow/Keras搭建的原生神经网络,默认没有添加sklearn要求的回归器标记,接口逻辑也不匹配,无法通过校验。
可行解决方案

方案1:封装模型适配Yellowbrick要求

针对statsmodels OLS模型

使用statsmodels官方提供的sklearn兼容适配器SMWrapper封装OLS模型,封装后会自动匹配sklearn接口规范:

from statsmodels.tools.sklearn import SMWrapper
import statsmodels.api as sm

# 封装OLS为sklearn兼容的回归器
model = SMWrapper(sm.OLS, fit_intercept=True)

# 后续即可正常传入PredictionError使用

如果不需要用到statsmodels的特殊统计功能,也可以直接替换为sklearn原生的LinearRegression,效果和sm.OLS完全一致,无需额外封装。

针对Keras神经网络模型

使用scikeras库的KerasRegressor封装原生Keras模型,自动适配sklearn接口要求:

from scikeras.wrappers import KerasRegressor

# 定义你的Keras模型结构
def build_model():
    from tensorflow.keras.models import Sequential
    from tensorflow.keras.layers import Dense
    model = Sequential()
    model.add(Dense(16, activation='relu', input_shape=(X_train.shape[1],)))
    model.add(Dense(1))
    model.compile(optimizer='adam', loss='mse')
    return model

# 封装为sklearn兼容的回归器
model = KerasRegressor(model=build_model, epochs=20, batch_size=16, verbose=0)

# 后续即可正常传入PredictionError使用

方案2:手动实现预测误差图(无需修改模型)

如果不想对原有模型做改造,可以直接手动绘制和Yellowbrick逻辑一致的预测误差图:

import matplotlib.pyplot as plt
import numpy as np

# 计算预测值
y_train_pred = model.predict(X_train)
y_test_pred = model.predict(X_test)

# 合并所有真实值获取坐标范围
all_y = np.concatenate([y_train, y_test])
min_val, max_val = all_y.min(), all_y.max()

# 绘制图形
plt.figure(figsize=(8,6))
plt.scatter(y_train, y_train_pred, alpha=0.6, label='训练集')
plt.scatter(y_test, y_test_pred, alpha=0.6, label='测试集')
plt.plot([min_val, max_val], [min_val, max_val], 'r--', label='完美预测线')

plt.xlabel('真实值')
plt.ylabel('预测值')
plt.title('预测误差图')
plt.legend()
plt.grid(alpha=0.3)
plt.show()

内容的提问来源于stack exchange,提问作者Joehat

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.06 15:39:03