CatBoost randomized_search如何传入指定训练/测试拆分的cv参数
CatBoost randomized_search 固定训练/验证集传参方法
根据CatBoost官方定义,randomized_search的cv参数支持传入预定义的数据集拆分,跳过默认的交叉验证逻辑;该传参方式要求传入可生成(训练索引数组, 测试索引数组)元组的可迭代对象。
原有写法的错误点
之前的传参cv={train_index,test_index}不符合要求,核心问题有两个:
- Python集合是无序结构,方法无法区分传入的两个索引哪个对应训练集、哪个对应测试集
- 哪怕把集合换成裸元组
(train_index, test_index)也不符合格式要求,方法会将其识别为多组拆分的迭代序列,而非单组固定拆分。
正确写法
哪怕只使用1组固定训练/验证拆分,也需要将(训练索引, 测试索引)的元组包裹在列表(最常用的可迭代结构)中传入。以10行数据、前5行训练后5行验证的场景为例,正确代码如下:
import numpy as np # 提取对应索引,转成numpy数组可避免pandas索引类型兼容问题 train_index = np.array(X[0:5].index) test_index = np.array(X[5:10].index) a_search = model.randomized_search( param_distributions=params, X=X, y=y, n_iter=5, # 核心:单组拆分元组放在列表中传入,完全跳过交叉验证 cv=[(train_index, test_index)] )
补充说明
- 按上述方式传参后,参数搜索全程只会使用指定的固定行范围做训练和验证,不会执行多折数据切分,可显著提升训练速度
- 如果需要传入多组固定拆分,只需在cv对应的列表中追加拆分元组即可,格式为
cv=[(train1, test1), (train2, test2)] - 传入的索引需要和X、y的行索引匹配,支持整数位置索引、DataFrame标签索引,不支持直接传入切片对象。
内容的提问来源于stack exchange,提问作者JPErwin
相关产品推荐
相关产品推荐

