如何用statsmodels.graphics绘制含交互项的偏回归图?
针对statsmodels交互项偏回归图的解决方案
statsmodels.graphics可以实现交互项的偏回归图绘制,针对你遇到的两个问题,具体解决方法如下:
问题1:plot_partregress无法传入拟合结果与指定参考水平
plot_partregress不支持直接传入拟合后的模型对象,但可以基于模型生成的设计矩阵手动绘制,同时保留公式中设置的分类变量参考水平:
- 先完成模型拟合:
results = reg.fit()
- 提取模型的设计矩阵与因变量(已包含交互项和按
Treatment编码的分类变量):
X = results.model.exog y = results.model.endog
- 确认交互项的列名(通过
results.model.exog_names查看),然后调用plot_partregress绘制:
from statsmodels.graphics.regressionplots import plot_partregress # 替换为你的交互项列名,比如'weight_del6_relative:pc1' interact_col = 'weight_del6_relative:pc1' col_idx = results.model.exog_names.index(interact_col) plot_partregress( endog=y, exog_idx=col_idx, exog=X, exog_names=results.model.exog_names, title=f"偏回归图:{interact_col}" )
问题2:从子图中单独提取目标偏回归图
plot_regress_exog和plot_partregress_grid返回的是matplotlib Figure对象,可通过axes属性定位并提取单个子图:
方法1:从plot_regress_exog中提取
from statsmodels.graphics.regressionplots import plot_regress_exog import matplotlib.pyplot as plt # 生成交互项的相关图组 fig = plot_regress_exog(results, 'weight_del6_relative:pc1') # 偏回归图通常是第2个子图(索引为1,可根据实际输出调整) ax_partreg = fig.axes[1] # 单独保存子图 ax_partreg.figure.savefig('交互项偏回归图.png') # 单独显示子图 plt.show(ax_partreg.figure)
方法2:从plot_partregress_grid中提取
from statsmodels.graphics.regressionplots import plot_partregress_grid import matplotlib.pyplot as plt fig = plot_partregress_grid(results) # 找到交互项对应的子图索引 interact_idx = results.model.exog_names.index('weight_del6_relative:pc1') ax_target = fig.axes[interact_idx] # 自定义子图样式并保存 ax_target.set_title("交互项偏回归图") ax_target.figure.savefig('单张偏回归图.png')
额外方案:手动绘制自定义偏回归图
如果需要更高自由度,可手动计算偏残差与偏拟合值,直接用matplotlib绘制:
import matplotlib.pyplot as plt import numpy as np # 获取交互项之外的所有变量参数与列 other_params = np.delete(results.params, col_idx) other_X = np.delete(X, col_idx, axis=1) # 计算偏残差(y减去其他变量的预测值) partial_resid = y - other_X @ other_params # 计算偏拟合值(交互项变量的拟合部分) partial_fit = X[:, col_idx] * results.params[col_idx] # 绘制散点图+拟合线 plt.scatter(X[:, col_idx], partial_resid, alpha=0.5) plt.plot(X[:, col_idx], partial_fit, color='red') plt.xlabel(interact_col) plt.ylabel("偏残差") plt.title("交互项偏回归图") plt.show()
内容的提问来源于stack exchange,提问作者Luis_12345
相关产品推荐
相关产品推荐

