PySpark 到 Scala:带有 StructType、GenericRowWithSchema 的 UDF 不能转换为 org.apache.spark.sql.Column

nev*_*_me 3 scala apache-spark pyspark

我有一些用 PySpark 编写的代码,我正忙着将它转换为 Scala。它一直进展顺利,除了现在我在 Scala 中为用户定义的函数而苦苦挣扎。

Python

from pyspark.sql import SparkSession
from pyspark.sql import SQLContext
from pyspark.sql.types import *
from pyspark.sql import functions as F

spark = SparkSession.builder.master('local[*]').getOrCreate()

a = spark.sparkContext.parallelize([(1,), (2,), (3,), (4,), (5,), (6,), (7,), (8,), (9,), (10,)]).toDF(["index"]).withColumn("a1", F.lit(1)).withColumn("a2", F.lit(2)).withColumn("a3", F.lit(3))

a = a.select("index", F.struct(*('a' + str(c) for c in range(1, 4))).alias('a'))

a.show()

def a_to_b(a):
    # 1. check if a technical cure exists
    b = {}
    for i in range(1, 4):
        b.update({'b' + str(i): a[i - 1] ** 2})
    return b

a_to_b_udf = F.udf(lambda x: a_to_b(x), StructType(list(StructField("b" + str(x), IntegerType()) for x in range(1, 4))))

b = a.select("index", "a", a_to_b_udf(a.a).alias("b"))

b.show()
Run Code Online (Sandbox Code Playgroud)

这产生:

+-----+-------+
|index|      a|
+-----+-------+
|    1|[1,2,3]|
|    2|[1,2,3]|
|    3|[1,2,3]|
|    4|[1,2,3]|
|    5|[1,2,3]|
|    6|[1,2,3]|
|    7|[1,2,3]|
|    8|[1,2,3]|
|    9|[1,2,3]|
|   10|[1,2,3]|
+-----+-------+
Run Code Online (Sandbox Code Playgroud)

和

+-----+-------+-------+
|index|      a|      b|
+-----+-------+-------+
|    1|[1,2,3]|[1,4,9]|
|    2|[1,2,3]|[1,4,9]|
|    3|[1,2,3]|[1,4,9]|
|    4|[1,2,3]|[1,4,9]|
|    5|[1,2,3]|[1,4,9]|
|    6|[1,2,3]|[1,4,9]|
|    7|[1,2,3]|[1,4,9]|
|    8|[1,2,3]|[1,4,9]|
|    9|[1,2,3]|[1,4,9]|
|   10|[1,2,3]|[1,4,9]|
+-----+-------+-------+
Run Code Online (Sandbox Code Playgroud)

斯卡拉

import org.apache.spark.sql._
import org.apache.spark.sql.functions._

// can ignore if running on spark-shell
val spark: SparkSession = SparkSession.builder()
  .master("local[*]")
  .getOrCreate()

import spark.implicits._

var a = spark.sparkContext.parallelize(1 to 10).toDF("index").withColumn("a1", lit(1)).withColumn("a2", lit(2)).withColumn("a3", lit(3))

// convert a{x} to struct column
a = a.select($"index", struct((1 to 3).map {x => col("a" + x)}.toList:_*).alias("a"))

a.show()

// this is where I am struggling, I have tried supplying a schema, but still get errors
val f = udf((a: Column) => {
  Seq(Math.pow(a(0).asInstanceOf[Double], 2), Math.pow(a(1).asInstanceOf[Double], 2), Math.pow(a(2).asInstanceOf[Double], 2))
})

val b = a.select($"index", $"a", f($"a").alias("b"))

// throws the below error
b.show()
Run Code Online (Sandbox Code Playgroud)

我可以显示()第一个数据帧,但是在尝试显示b.

错误是:

java.lang.ClassCastException: org.apache.spark.sql.catalyst.expressions.GenericRowWithSchema cannot be cast to org.apache.spark.sql.Column
  at $line23.$read$$iw$$iw$$iw$$iw$$iw$$iw$$iw$$iw$$iw$$iw$$iw$$iw$$anonfun$1.apply(<console>:31)
  at org.apache.spark.sql.catalyst.expressions.GeneratedClass$GeneratedIterator.processNext(Unknown Source)
  at org.apache.spark.sql.execution.BufferedRowIterator.hasNext(BufferedRowIterator.java:43)
  at org.apache.spark.sql.execution.WholeStageCodegenExec$$anonfun$8$$anon$1.hasNext(WholeStageCodegenExec.scala:370)
  at org.apache.spark.sql.execution.SparkPlan$$anonfun$4.apply(SparkPlan.scala:246)
  at org.apache.spark.sql.execution.SparkPlan$$anonfun$4.apply(SparkPlan.scala:240)
  at org.apache.spark.rdd.RDD$$anonfun$mapPartitionsInternal$1$$anonfun$apply$24.apply(RDD.scala:784)
  at org.apache.spark.rdd.RDD$$anonfun$mapPartitionsInternal$1$$anonfun$apply$24.apply(RDD.scala:784)
  at org.apache.spark.rdd.MapPartitionsRDD.compute(MapPartitionsRDD.scala:38)
  at org.apache.spark.rdd.RDD.computeOrReadCheckpoint(RDD.scala:319)
  at org.apache.spark.rdd.RDD.iterator(RDD.scala:283)
  at org.apache.spark.scheduler.ResultTask.runTask(ResultTask.scala:70)
  at org.apache.spark.scheduler.Task.run(Task.scala:85)
  at org.apache.spark.executor.Executor$TaskRunner.run(Executor.scala:274)
  at java.util.concurrent.ThreadPoolExecutor.runWorker(ThreadPoolExecutor.java:1142)
  at java.util.concurrent.ThreadPoolExecutor$Worker.run(ThreadPoolExecutor.java:617)
  at java.lang.Thread.run(Thread.java:745)
Run Code Online (Sandbox Code Playgroud)

我已经尝试像在 Python 中所做的那样为我的 UDF 设置模式,但我仍然遇到相同的错误。

有谁知道我如何解决这个问题?我的例子很简单,但我需要对 UDF 做的是在返回结构体之前进行相当多的转换。

nev*_*_me 5

我觉得很傻,因为我从周五下午开始就一直在挣扎。

从具有复杂输入参数的 Spark Sql UDF,

结构类型转换为 org.apache.spark.sql.Row

我的问题是Column我提供给我的函数的类型。

val f = udf((a: Column) => {
  Seq(Math.pow(a(0).asInstanceOf[Double], 2), Math.pow(a(1).asInstanceOf[Double], 2), Math.pow(a(2).asInstanceOf[Double], 2))
})
Run Code Online (Sandbox Code Playgroud)

我应该用它Row来代替。

val f = udf((a: Row) => {
  println("testing")
  Seq(Math.pow(a(0).asInstanceOf[Int], 2).asInstanceOf[Int],
    Math.pow(a(1).asInstanceOf[Int], 2).asInstanceOf[Int],
    Math.pow(a(2).asInstanceOf[Int], 2).asInstanceOf[Int])
})
Run Code Online (Sandbox Code Playgroud)