使用statsmodels进行条件logit模型概率预测时遇NotImplementedError
解决ConditionalLogit predict() NotImplementedError问题
问题根源
statsmodels 0.12.2版本的ConditionalLogit类并未实现predict()方法,官方文档描述可能滞后于实际功能进度,这就是调用报错的直接原因。
方案一:手动计算预测概率
条件logistic回归核心是基于组内相对概率建模,我们可以用模型参数手动推导预测值,步骤如下:
- 计算线性预测值:
xb = 自变量 × 模型参数(若自变量是多维数组,改用矩阵乘法x @ params) - 对每个分组,将组内样本的
xb指数化得到exp(xb) - 每个样本的预测概率 = 自身
exp(xb)÷ 组内所有exp(xb)的总和
代码示例
from statsmodels.discrete.conditional_models import ConditionalLogit import numpy as np # 生成模拟数据 g = np.kron(np.arange(100), np.ones(5)).astype(int) x = np.random.normal(size=500) pr = 1 / (1 + np.exp(-x)) y = (np.random.uniform(size=500) < pr).astype(int) # 拟合模型 m = ConditionalLogit(endog=y, exog=x, groups=g) r = m.fit() # 手动计算预测概率 params = r.params xb = x * params # 多维自变量请替换为 x @ params pred_probs = np.zeros_like(xb) unique_groups = np.unique(g) for group in unique_groups: mask = g == group group_exp_vals = np.exp(xb[mask]) pred_probs[mask] = group_exp_vals / group_exp_vals.sum() # 查看前5个样本的预测概率 print(pred_probs[:5])
方案二:升级statsmodels版本
statsmodels在0.13.0及后续版本中已补充实现ConditionalLogit.predict()方法,直接升级包即可解决问题:
pip install --upgrade statsmodels
升级后可像无条件Logit模型那样直接调用r.predict(x),内部逻辑与手动计算的组内归一化逻辑一致。
内容的提问来源于stack exchange,提问作者Ali Fahmi
相关产品推荐
相关产品推荐

