如何在PySpark中通过Py4J将方法作为参数传入Scala函数?
问题描述
我正在使用PySpark,希望通过Py4J将一个对象的方法作为参数传入函数调用。我有一个.jar库,其中包含一个签名如下的方法(省略其他参数细节):
def apply(name: String, exp: Column => Column): SomeClass = { ... }
该方法属于some.pkg包下SomeClass的伴生对象。将jar加载到PySpark后,我尝试通过以下代码调用该方法:
>>> spark._jvm.some.pkg.SomeClass.apply("", spark._jvm.org.apache.sql.functions.hour)
但出现了如下错误:
Traceback (most recent call last): File "<stdin>", line 1, in <module> File "/usr/local/lib/python3.9/site-packages/py4j/java_gateway.py", line 1296, in __call__ args_command, temp_args = self._build_args(*args) File "/usr/local/lib/python3.9/site-packages/py4j/java_gateway.py", line 1266, in _build_args [get_command_part(arg, self.pool) for arg in new_args]) File "/usr/local/lib/python3.9/site-packages/py4j/java_gateway.py", line 1266, in <listcomp> [get_command_part(arg, self.pool) for arg in new_args]) File "/usr/local/lib/python3.9/site-packages/py4j/protocol.py", line 298, in get_command_part command_part = REFERENCE_TYPE + parameter._get_object_id() AttributeError: 'JavaMember' object has no attribute '_get_object_id'
Py4J无法直接将方法作为参数传入其他函数,以下是可行的解决办法:
解决办法
方案1:用Scala匿名函数包装目标方法
Py4J无法直接传递Java/Scala方法引用,但可以通过spark._jvm.scala.runtime.AbstractFunction1创建包装目标方法的匿名函数对象,再传递给目标方法:# 创建包装hour方法的匿名函数 hour_func = spark._jvm.scala.runtime.AbstractFunction1(lambda col: spark._jvm.org.apache.sql.functions.hour(col)) # 调用目标方法 spark._jvm.some.pkg.SomeClass.apply("", hour_func)方案2:编写Scala辅助类
如果需要多次传递方法引用,建议编写Scala辅助类,提供静态方法返回包装好的函数对象,打包进jar后在PySpark中调用:package some.pkg import org.apache.spark.sql.Column import org.apache.spark.sql.functions._ object FunctionWrappers { def wrapHour: Column => Column = hour(_) // 可添加更多方法包装逻辑 }PySpark调用代码:
spark._jvm.some.pkg.SomeClass.apply("", spark._jvm.some.pkg.FunctionWrappers.wrapHour())方案3:使用Py4J的
JavaLambda(Py4J 0.10.9+适用)
若Py4J版本在0.10.9及以上,可直接用JavaLambda包装Python函数,Py4J会自动转换为对应Java函数接口:from py4j.java_gateway import JavaLambda def hour_wrapper(col): return spark._jvm.org.apache.sql.functions.hour(col) # 包装为JavaLambda,指定对应函数接口类型 java_hour_func = JavaLambda(hour_wrapper, "org.apache.spark.api.java.function.Function1") spark._jvm.some.pkg.SomeClass.apply("", java_hour_func)
内容的提问来源于stack exchange,提问作者PiFace
相关产品推荐
相关产品推荐

