【问题标题】:Deep decision tree in PySparkPySpark 中的深度决策树
【发布时间】:2018-09-21 11:45:00
【问题描述】:

我正在使用 PySpark 进行机器学习,我想训练决策树分类器、随机森林和梯度提升树。我想尝试不同的最大深度值,并通过网格搜索和交叉验证选择最好的一个。但是,Spark 告诉我 DecisionTree 目前仅支持 maxDepth

from pyspark.ml import Pipeline
from pyspark.ml.classification import RandomForestClassifier
from pyspark.ml.feature import IndexToString, StringIndexer, VectorIndexer
from pyspark.ml.evaluation import MulticlassClassificationEvaluator
from pyspark.ml.tuning import CrossValidator, ParamGridBuilder

# Load and parse the data file, converting it to a DataFrame.

data = spark.read.format("libsvm").load("data/mllib/sample_libsvm_data.txt")

 # Index labels, adding metadata to the label column.
 # Fit on whole dataset to include all labels in index.

 labelIndexer = StringIndexer(inputCol="label", 
outputCol="indexedLabel").fit(data)

 # Automatically identify categorical features, and index them.
 # Set maxCategories so features with > 4 distinct values are treated as continuous.
featureIndexer =\
VectorIndexer(inputCol="features", outputCol="indexedFeatures", maxCategories=4).fit(data)

# Split the data into training and test sets (30% held out for testing)
(trainingData, testData) = data.randomSplit([0.7, 0.3])

# Train a RandomForest model.

 rf = RandomForestClassifier(labelCol="indexedLabel", 
      featuresCol="indexedFeatures", numTrees=500)

 # Convert indexed labels back to original labels.
labelConverter = IndexToString(inputCol="prediction", 
outputCol="predictedLabel",
                           labels=labelIndexer.labels)

# Chain indexers and forest in a Pipeline
 pipeline = Pipeline(stages=[labelIndexer, featureIndexer, rf, labelConverter])

 paramGrid_rf = ParamGridBuilder() \
   .addGrid(rf.maxDepth, [50,100,150,250,300]) \
   .build()

 crossval_rf = CrossValidator(estimator=pipeline,
                       estimatorParamMaps=paramGrid_rf,
                      evaluator=BinaryClassificationEvaluator(),
                      numFolds= 5) 

 cvModel_rf = crossval_rf.fit(trainingData)

上面的代码给了我下面的错误信息。

Py4JJavaError:调用 o12383.fit 时出错。 :java.lang.IllegalArgumentException:要求失败:DecisionTree目前仅支持maxDepth

【问题讨论】:

    标签: pyspark


    【解决方案1】:
    猜你喜欢
    • 2017-03-07
    • 2016-10-18
    • 2019-03-10
    • 2017-01-21
    • 2023-04-08
    • 2017-03-12
    • 2017-10-05
    • 2017-09-17
    • 2022-01-27
    相关资源
    最近更新 更多