Gwy*_*nFR 5 python machine-learning hyperparameters pyspark apache-spark-ml
我正在使用PySpark 2.0进行Kaggle比赛。我想知道模型(RandomForest)的行为,具体取决于不同的参数。ParamGridBuilder()允许为单个参数指定不同的值,然后执行(我想)整个参数集的笛卡尔积。假设我DataFrame已经定义:
rdc = RandomForestClassifier()
pipeline = Pipeline(stages=STAGES + [rdc])
paramGrid = ParamGridBuilder().addGrid(rdc.maxDepth, [3, 10, 20])
.addGrid(rdc.minInfoGain, [0.01, 0.001])
.addGrid(rdc.numTrees, [5, 10, 20, 30])
.build()
evaluator = MulticlassClassificationEvaluator()
valid = TrainValidationSplit(estimator=pipeline,
estimatorParamMaps=paramGrid,
evaluator=evaluator,
trainRatio=0.50)
model = valid.fit(df)
result = model.bestModel.transform(df)
Run Code Online (Sandbox Code Playgroud)
好的,现在我可以使用手工功能检索简单的信息:
def evaluate(result):
predictionAndLabels = result.select("prediction", "label")
metrics = ["f1","weightedPrecision","weightedRecall","accuracy"]
for m in metrics:
evaluator = MulticlassClassificationEvaluator(metricName=m)
print(str(m) + ": " + str(evaluator.evaluate(predictionAndLabels)))
Run Code Online (Sandbox Code Playgroud)
现在我想要几件事:
print(model.validationMetrics)显示(似乎)包含每个模型准确性的列表,但是我不知道要引用哪个模型。如果我可以检索所有这些信息,则应该能够显示图形,条形图,并且可以像使用Panda和一样工作sklearn。
Spark 2.4+
SPARK- 21088 CrossValidator,TrainValidationSplit在拟合时应收集所有模型 -添加了对收集子模型的支持。
默认情况下,此行为是禁用的,但可以使用CollectSubModels Param(setCollectSubModels)进行控制。
valid = TrainValidationSplit(
estimator=pipeline,
estimatorParamMaps=paramGrid,
evaluator=evaluator,
collectSubModels=True)
model = valid.fit(df)
model.subModels
Run Code Online (Sandbox Code Playgroud)
火花<2.4
长话短说,您根本无法获取所有模型的参数,因为与相似CrossValidator,它TrainValidationSplitModel仅保留最佳模型。这些类是为半自动模型选择而不是探索或实验而设计的。
所有型号的参数是什么?
尽管您无法检索validationMetrics与输入对应的实际模型,Params所以您应该能够简单地将zip两者:
from typing import Dict, Tuple, List, Any
from pyspark.ml.param import Param
from pyspark.ml.tuning import TrainValidationSplitModel
EvalParam = List[Tuple[float, Dict[Param, Any]]]
def get_metrics_and_params(model: TrainValidationSplitModel) -> EvalParam:
return list(zip(model.validationMetrics, model.getEstimatorParamMaps()))
Run Code Online (Sandbox Code Playgroud)
得到一些有关指标和参数之间关系的信息。
如果您需要更多信息,则应使用PipelineParams。它将保留所有可用于进一步处理的模型:
models = pipeline.fit(df, params=paramGrid)
Run Code Online (Sandbox Code Playgroud)
它将生成PipelineModels与params参数相对应的列表:
zip(models, params)
Run Code Online (Sandbox Code Playgroud)
| 归档时间: |
|
| 查看次数: |
4571 次 |
| 最近记录: |