网格搜索

使用 PySpark 进行机器学习

Andrew Collier

Data Scientist, Fathom Data

选择最优参数值

使用 PySpark 进行机器学习

再次回到汽车数据

cars.select('mass', 'cyl', 'consumption').show(5)
+------+---+-----------+
|  mass|cyl|consumption|
+------+---+-----------+
|1451.0|  6|       9.05|
|1129.0|  4|       6.53|
|1399.0|  4|       7.84|
|1147.0|  4|       7.84|
|1111.0|  4|       9.05|
+------+---+-----------+
使用 PySpark 进行机器学习

含截距的油耗模型

带截距的线性回归。拟合训练数据。

regression = LinearRegression(labelCol='consumption', fitIntercept=True)
regression = regression.fit(cars_train)

在测试数据上计算 RMSE。

evaluator.evaluate(regression.transform(cars_test))
# 带截距模型的 RMSE
0.745974203928479
使用 PySpark 进行机器学习

不含截距的油耗模型

不含截距的线性回归。拟合训练数据。

regression = LinearRegression(labelCol='consumption', fitIntercept=False)
regression = regression.fit(cars_train)

在测试数据上计算 RMSE。

# 无截距模型(第二个模型)的 RMSE
0.852819012439
# 含截距模型(第一个模型)的 RMSE
0.745974203928
使用 PySpark 进行机器学习

参数网格

from pyspark.ml.tuning import ParamGridBuilder

# 创建参数网格构造器
params = ParamGridBuilder()

# 添加网格点 params = params.addGrid(regression.fitIntercept, [True, False])
# 构建网格 params = params.build()
# 有多少个模型? print('Number of models to be tested: ', len(params))
Number of models to be tested:  2
使用 PySpark 进行机器学习

带交叉验证的网格搜索

创建交叉验证器并拟合训练数据。

cv = CrossValidator(estimator=regression,
                    estimatorParamMaps=params,
                    evaluator=evaluator)
cv = cv.setNumFolds(10).setSeed(13).fit(cars_train)

每个模型的交叉验证 RMSE 是多少?

cv.avgMetrics
[0.800663722151, 0.907977823182]
使用 PySpark 进行机器学习

最佳模型与参数

# 访问最佳模型
cv.bestModel

也可直接使用交叉验证器对象。

predictions = cv.transform(cars_test)

获取最佳参数。

cv.bestModel.explainParam('fitIntercept')
'fitIntercept: whether to fit an intercept term (default: True, current: True)'
使用 PySpark 进行机器学习

更复杂的网格

params = ParamGridBuilder() \
            .addGrid(regression.fitIntercept, [True, False]) \

.addGrid(regression.regParam, [0.001, 0.01, 0.1, 1, 10]) \
.addGrid(regression.elasticNetParam, [0, 0.25, 0.5, 0.75, 1]) \ .build()

现在有多少个模型?

print ('Number of models to be tested: ', len(params))
Number of models to be tested:  50
使用 PySpark 进行机器学习

寻找最佳参数!

使用 PySpark 进行机器学习

Preparing Video For Download...