如何在PySpark Pipeline中使用Scala编写的自定义Transformer
如何在PySpark Pipeline中使用Scala编写的自定义Transformer?
我来帮你一步步搞定这个问题!把Scala写的UpperTransformer集成到PySpark Pipeline里其实没那么复杂,主要分为打包Scala代码、在PySpark中加载类、集成到Pipeline这几个核心步骤,下面详细说明:
第一步:打包Scala自定义Transformer为JAR包
首先得把你提供的Scala代码编译打包成JAR文件,这是跨语言调用的基础。
1. 完善Scala代码(添加包名)
为了在PySpark里能准确找到这个类,记得给你的Scala代码加上包名(比如com.example),修改后的代码如下:
package com.example import org.apache.spark.ml.UnaryTransformer import org.apache.spark.ml.util.Identifiable import org.apache.spark.sql.types.{DataType, StringType} class UpperTransformer(override val uid: String) extends UnaryTransformer[String, String, UpperTransformer] { def this() = this(Identifiable.randomUID("upper")) override protected def validateInputType(inputType: DataType): Unit = { require(inputType == StringType) } protected def createTransformFunc: String => String = { _.toUpperCase } protected def outputDataType: DataType = StringType }
2. 用SBT或Maven打包
这里以SBT为例,创建一个build.sbt文件,注意Scala版本要和你的PySpark版本匹配(比如PySpark 3.3.x默认用Scala 2.12):
name := "UpperTransformer" version := "1.0" scalaVersion := "2.12.15" // 对应PySpark的Scala版本,必须一致! libraryDependencies += "org.apache.spark" %% "spark-core" % "3.3.0" % Provided libraryDependencies += "org.apache.spark" %% "spark-sql" % "3.3.0" % Provided
然后在项目根目录运行sbt package,打包完成后,JAR文件会生成在target/scala-2.12/uppertransformer_2.12-1.0.jar路径下。
第二步:在PySpark中加载Scala Transformer
启动PySpark的时候,需要把刚才生成的JAR包引入,或者在SparkSession配置里指定。
方式1:启动PySpark时带JAR参数
pyspark --jars /path/to/uppertransformer_2.12-1.0.jar
方式2:在代码中配置SparkSession
from pyspark.sql import SparkSession spark = SparkSession.builder \ .appName("ScalaTransformerInPySpark") \ .config("spark.jars", "/path/to/uppertransformer_2.12-1.0.jar") \ .getOrCreate()
接下来,通过Spark的JVM桥接获取Scala的UpperTransformer类,并实例化它:
# 注意替换成你实际的包名! UpperTransformer = spark._jvm.com.example.UpperTransformer # 实例化Transformer,和Scala里的无参构造对应 upper_transformer = UpperTransformer()
第三步:集成到PySpark Pipeline中
现在你就可以像使用PySpark自带的Transformer一样,把upper_transformer加入Pipeline里了。举个完整的例子:
from pyspark.ml import Pipeline from pyspark.ml.feature import StringIndexer # 创建测试数据 test_data = spark.createDataFrame( [(1, "hello spark"), (2, "scala & pyspark"), (3, "custom transformer")], ["id", "original_text"] ) # 构建Pipeline,混合Scala和PySpark的Transformer pipeline = Pipeline(stages=[ # 设置输入输出列,和PySpark的API用法一致 upper_transformer.setInputCol("original_text").setOutputCol("upper_text"), # 再加一个PySpark自带的Feature处理stage StringIndexer(inputCol="upper_text", outputCol="text_index") ]) # 训练Pipeline并转换数据 pipeline_model = pipeline.fit(test_data) result_df = pipeline_model.transform(test_data) # 查看结果 result_df.select("original_text", "upper_text", "text_index").show(truncate=False)
运行后你会看到original_text列被转成大写的upper_text,同时生成了对应的索引列。
关键注意事项
- Scala版本必须匹配:PySpark和Scala的版本要严格对应,否则会出现类加载错误。比如PySpark 3.4.x对应Scala 2.12,PySpark 3.5.x开始支持Scala 2.13。
- 包名不能漏:Scala类必须带包名,否则在PySpark里无法通过
spark._jvm准确找到类。 - 参数设置兼容:如果你的Transformer有自定义参数,也可以通过
setXXX()方法设置,和PySpark的API用法完全一致,因为Spark ML是跨语言设计的。
内容的提问来源于stack exchange,提问作者pratyush
相关产品推荐
相关产品推荐

