使用yellowbrick绘制回归预测误差图时OLS、Keras模型报错如何解决?
报错原因
Yellowbrick的回归类可视化器完全基于scikit-learn的API规范设计,会调用sklearn内置的is_regressor()函数校验传入模型的类型,只有符合以下两个条件才会被判定为合法回归器:
- 模型带有
_estimator_type = 'regressor'的属性标记 - 模型遵循sklearn估计器的统一接口(有
fit、predict等标准方法,调用逻辑与sklearn一致)
你用到的三个模型中,RandomForestRegressor是sklearn原生实现,天然符合上述要求,所以可以正常出图;另外两个模型不满足规范,因此触发报错:
- statsmodels的
sm.OLS()实现的线性回归,没有遵循sklearn的接口设计,既没有内置_estimator_type标记,拟合、预测的调用逻辑也和sklearn不匹配,无法通过校验。 - 基于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
相关产品推荐
相关产品推荐

