使用sklearn的CategoricalNB调用predict()时报IndexError的解决求助
问题背景
我在使用Python的sklearn.naive_bayes的CategoricalNB()模型时,训练过程无报错,但调用predict()方法时抛出如下错误:
C:\ProgramData\Anaconda3\envs\Python\lib\site-packages\sklearn\naive_bayes.py in predict(self, X) 81 check_is_fitted(self) 82 X = self._check_X(X) ---> 83 jll = self._joint_log_likelihood(X) 84 return self.classes_[np.argmax(jll, axis=1)] 85 C:\ProgramData\Anaconda3\envs\Python\lib\site-packages\sklearn\naive_bayes.py in _joint_log_likelihood(self, X) 1459 for i in range(self.n_features_in_): 1460 indices = X[:, i] -> 1461 jll += self.feature_log_prob_[i][:, indices].T 1462 total_ll = jll + self.class_log_prior_ 1463 return total_ll IndexError: index 6 is out of bounds for axis 1 with size 6
报错出现在以下代码的最后一行:
model = CategoricalNB() model.fit(X_train, y_train) y_train_pred = model.predict(X_train) y_test_pred = model.predict(X_test)
相关数据参数:X_train形状为(1318, 12),类型为numpy.ndarray;y_train形状为(1318,),类型为pandas.core.series.Series;X_test形状为(566, 12),类型为numpy.ndarray。输入特征已通过OrdinalEncoder()编码,目标变量通过LabelEncoder()编码,每列取值为0到对应类别数的正整数,目标为0到6的多分类任务。
问题更新
更新1
我在train_test_split()中添加stratify=y参数后问题暂时解决,修改后代码如下:
from sklearn.model_selection import train_test_split X_train, X_test, y_train, y_test = train_test_split(X, y, test_size = 0.3, stratify=y, random_state = 123)
但我不理解该参数为何能解决索引问题,我对stratify的认知仅为保证训练集、测试集各类别比例与原数据集一致。
更新2
我的数据集包含多个输出变量,上述方案仅对部分输出有效,部分输出仍会抛出IndexError。
更新3
我查询后发现该问题属于已知bug,原因是测试集出现了训练集中未覆盖的类别。请问有什么方案可以绕过该bug?
原因说明
- 报错本质:CategoricalNB训练时会按照训练集每个特征实际出现的类别数量生成概率表,若测试集中某个特征出现了训练集完全没有的编码值,就会直接触发索引越界。
- stratify临时生效的原因:该参数仅保证标签的分布和原数据集一致,一定程度上降低了和标签强相关的稀有特征类别被全部分到测试集的概率,但无法解决特征维度的稀有类别漏采问题,因此只对部分场景有效。
绕过方案
- 方案1:编码阶段提前处理未知类别。使用OrdinalEncoder时指定
handle_unknown='use_encoded_value'和unknown_value参数,将训练集未出现的特征值统一映射为预设的整数值,同时给CategoricalNB设置min_categories参数,给每个特征预留未知类别的概率表位置。示例代码如下:
# 假设每个特征最多有7个类别,未知值编码为6 encoder = OrdinalEncoder(handle_unknown='use_encoded_value', unknown_value=6) # 给CategoricalNB指定每个特征的最小类别数,预留未知类位置 model = CategoricalNB(min_categories=7)
- 方案2:训练前统一处理特征类别。遍历所有特征列统计全局所有可能取值,切分数据集后给训练集补入每个特征的全部取值对应的少量样本,保证训练集覆盖所有特征的全部类别,避免测试集出现训练集未见过的特征值。
- 方案3:自定义模型逻辑。继承CategoricalNB重写
_joint_log_likelihood方法,遇到超出索引范围的特征值时,直接给该特征对应的类别概率赋极小值(如对数化的1e-9),规避索引报错。 - 方案4:特征预处理阶段合并稀有类别。将每个特征中出现频次极低的类别统一合并为“其他”类别后再编码训练,从根源上避免测试集出现训练集未覆盖的特征类别。
内容的提问来源于stack exchange,提问作者froot.cocktail
相关产品推荐
相关产品推荐

