能否在不提供参数网格的情况下使用Spark MLLib的CrossValidator?
问题原因
你遇到的报错确实是未传入estimatorParamMaps参数导致的。Spark ML库的CrossValidator组件最初设计为超参数调优工具,因此estimatorParamMaps(待搜索的参数映射列表)是必填参数,没有默认值。
无参数网格的交叉验证实现
如果你不需要做超参数搜索,只想用交叉验证评估固定参数的模型性能,只需要构造一个仅包含单组默认参数的参数映射列表即可,实现方式如下:
from pyspark.ml.tuning import ParamGridBuilder, CrossValidator # 构造无搜索项的参数网格,仅保留模型当前的设置参数 param_grid = ParamGridBuilder().build() # 初始化交叉验证器,传入构造好的参数网格 cv = CrossValidator( estimator=randomForestRegressor, evaluator=evaluator, numFolds=5, estimatorParamMaps=param_grid ) cv_model = cv.fit(vsdf)
结果说明
- 上述代码中
ParamGridBuilder().build()会生成一个仅包含空参数映射的列表,CrossValidator会在每折训练时直接使用你定义的randomForestRegressor的原有参数,不会执行额外的超参数搜索。如果你的随机森林回归器已经自定义了部分参数(比如树的数量、最大深度等),该方法依然有效。 - 训练完成后
cv_model.avgMetrics中会返回唯一的评估值,就是5折交叉验证得到的平均模型性能,符合你评估测试错误率的需求。
内容的提问来源于stack exchange,提问作者Anirban Chakraborty
相关产品推荐
相关产品推荐

