带MultiIndex的DataFrame向LightGBM交叉验证传自定义折叠索引报错
问题分析与解决
核心问题
LightGBM的lgb.cv的folds参数不接受Pandas的标签索引(包括MultiIndex的元组标签),它要求传入数据的整数位置索引,或者可被转换为一维numpy数组/Series的整数序列。你遇到的错误是因为传入了MultiIndex的标签元组,LightGBM无法将其解析为有效的位置索引。
解决方法
将MultiIndex的标签转换为对应的整数位置下标,可通过Pandas的Index.get_indexer()方法实现。
修改后的可运行代码
import lightgbm as lgb import pandas as pd df = pd.DataFrame({ "index_1": ["i1", "i2"], "index_2": ["j1", "j2"], "feature": [1, 2], "target": [3, 4] }).set_index(["index_1", "index_2"]) data = lgb.Dataset(data=df["feature"].to_frame(), label=df["target"]) # 将MultiIndex标签转换为整数位置下标 train_labels = [("i1", "j1")] test_labels = [("i2", "j2")] train_idx = df.index.get_indexer(train_labels) test_idx = df.index.get_indexer(test_labels) # 构造符合要求的folds结构 folds = [ (train_idx, test_idx), (train_idx, test_idx) ] gbm = lgb.cv( params={"objective": "regression"}, # 补充必要的任务类型参数 train_set=data, folds=folds, ) print(gbm)
关键说明
- LightGBM的Dataset内部基于数据的整数位置管理样本,不会保留原始DataFrame的行索引(包括MultiIndex)信息,因此无法直接识别标签索引。
folds参数的每个元素必须是(训练集位置索引序列,测试集位置索引序列)的二元组,序列类型可以是列表、一维numpy数组或Pandas Series,元素必须为整数。
内容的提问来源于stack exchange,提问作者GlaceCelery
相关产品推荐
相关产品推荐

