移除测试集新类别:基于Sklearn的时间序列交叉验证实现问询
时间序列交叉验证中强制限制特征类别范围的解决方案
不用从零构建交叉验证流程,基于Sklearn现有组件做小幅修改就能实现需求。核心思路是给模型加类别校验包装器,或自定义带类别检查的时间拆分器,确保测试集的特征A、B类别仅包含训练阶段见过的取值,否则触发报错。
方案一:自定义模型包装器(推荐)
给你的基础模型套一层包装,在训练时记录特征A、B的所有已见类别,预测时强制校验测试集的类别范围,不符合则抛出错误。这种方式完全兼容Sklearn的cross_validate和TimeSeriesSplit。
代码实现
import numpy as np from sklearn.base import BaseEstimator, ClassifierMixin, clone from sklearn.model_selection import TimeSeriesSplit, cross_validate from sklearn.ensemble import RandomForestClassifier class CategoryRestrictedModel(BaseEstimator, ClassifierMixin): def __init__(self, base_model, cat_features=['A', 'B']): self.base_model = base_model self.cat_features = cat_features self.seen_categories = {} # 存储每个特征的已训练类别集合 def fit(self, X, y=None): # 克隆并训练基础模型 self.model_ = clone(self.base_model) self.model_.fit(X, y) # 记录训练集中指定特征的所有类别 for feat in self.cat_features: self.seen_categories[feat] = set(X[feat].unique()) return self def predict(self, X): # 校验测试集的特征类别 for feat in self.cat_features: unseen_cats = set(X[feat].unique()) - self.seen_categories[feat] if unseen_cats: raise ValueError(f"特征[{feat}]出现训练阶段未见过的类别: {unseen_cats}") # 执行基础模型的预测 return self.model_.predict(X) # 可选:若需要概率预测,同样包装并校验 def predict_proba(self, X): for feat in self.cat_features: unseen_cats = set(X[feat].unique()) - self.seen_categories[feat] if unseen_cats: raise ValueError(f"特征[{feat}]出现训练阶段未见过的类别: {unseen_cats}") return self.model_.predict_proba(X)
使用方式
# 假设X是包含特征A、B的DataFrame,y是目标变量 # 初始化基础模型(这里用随机森林示例,可替换为你的模型) base_model = RandomForestClassifier() # 初始化带类别限制的包装模型 restricted_model = CategoryRestrictedModel(base_model, cat_features=['A', 'B']) # 设置时间序列交叉验证拆分 tscv = TimeSeriesSplit(n_splits=5) # 运行交叉验证,若测试集出现未见过的类别会直接报错 cv_results = cross_validate(restricted_model, X, y, cv=tscv, scoring=['accuracy'])
方案二:自定义带类别检查的时间拆分器
如果希望在交叉验证的数据拆分阶段就提前检查测试集的类别(而非等到预测时),可以继承TimeSeriesSplit,在生成训练/测试索引后立即校验类别范围。
代码实现
class CategoryCheckedTimeSeriesSplit(TimeSeriesSplit): def __init__(self, n_splits=5, cat_features=['A', 'B']): super().__init__(n_splits=n_splits) self.cat_features = cat_features def split(self, X, y=None, groups=None): for train_idx, test_idx in super().split(X, y, groups): X_train = X.iloc[train_idx] X_test = X.iloc[test_idx] # 检查每个指定特征的测试集类别是否全部在训练集中存在 for feat in self.cat_features: train_cats = set(X_train[feat].unique()) test_cats = set(X_test[feat].unique()) unseen_cats = test_cats - train_cats if unseen_cats: raise ValueError(f"当前Fold测试集特征[{feat}]出现未见过的类别: {unseen_cats}") yield train_idx, test_idx
使用方式
# 初始化带类别检查的时间拆分器 tscv_checked = CategoryCheckedTimeSeriesSplit(n_splits=5, cat_features=['A', 'B']) # 运行交叉验证,拆分阶段发现问题直接报错 cv_results = cross_validate(restricted_model, X, y, cv=tscv_checked, scoring=['accuracy'])
两种方案的选择
- 方案一适合需要在预测环节严格校验的场景,确保模型在任何预测场景下都不会输出未见过的类别;
- 方案二适合需要提前排查数据拆分问题的场景,避免无效的模型训练过程。
内容的提问来源于stack exchange,提问作者Slevin_42
相关产品推荐
相关产品推荐

