【发布时间】:2016-06-24 15:21:58
【问题描述】:
我有一个 id(标签)范围从 1 到 1040 的数据库。我正在使用多类 Logistic 回归来预测 id。现在,如果我只想训练标签的一个子集,比如说从 800 到 810。当我为 11 个类设置 setNumClasses(11) 时出现错误。我必须始终将此方法设置为类的最大值,即 1040。这样训练模型将针对从 0 到 1040 的所有标签进行训练,这非常昂贵并且使用大量资源。
我理解的对吗?如何通过给定 setNumClasses(count_of_classes) 仅针对标签子集训练我的模型。
final LogisticRegressionModel model = new LogisticRegressionWithLBFGS()
.setNumClasses(811).run(train.rdd());
【问题讨论】:
标签: apache-spark logistic-regression apache-spark-mllib