【发布时间】:2017-01-22 14:50:45
【问题描述】:
以下 Spark 代码正确演示了我想要做什么,并使用一个很小的演示数据集生成正确的输出。
当我在大量生产数据上运行相同类型的通用代码时,我遇到了运行时问题。 Spark 作业在我的集群上运行了大约 12 个小时,但失败了。
只看下面的代码,将每一行都分解掉似乎效率很低,只是将其合并回来。在给定的测试数据集中,array_value_1 中包含三个值和array_value_2 中的三个值的第四行将爆炸为 3*3 或九个爆炸行。
那么,在一个更大的数据集中,一行有五个这样的数组列,每列有十个值,会爆炸成 10^5 行吗?
查看提供的 Spark 函数,没有开箱即用的函数可以满足我的需求。我可以提供一个用户定义的函数。这样做有速度上的缺点吗?
val sparkSession = SparkSession.builder.
master("local")
.appName("merge list test")
.getOrCreate()
val schema = StructType(
StructField("category", IntegerType) ::
StructField("array_value_1", ArrayType(StringType)) ::
StructField("array_value_2", ArrayType(StringType)) ::
Nil)
val rows = List(
Row(1, List("a", "b"), List("u", "v")),
Row(1, List("b", "c"), List("v", "w")),
Row(2, List("c", "d"), List("w")),
Row(2, List("c", "d", "e"), List("x", "y", "z"))
)
val df = sparkSession.createDataFrame(rows.asJava, schema)
val dfExploded = df.
withColumn("scalar_1", explode(col("array_value_1"))).
withColumn("scalar_2", explode(col("array_value_2")))
// This will output 19. 2*2 + 2*2 + 2*1 + 3*3 = 19
logger.info(s"dfExploded.count()=${dfExploded.count()}")
val dfOutput = dfExploded.groupBy("category").agg(
collect_set("scalar_1").alias("combined_values_2"),
collect_set("scalar_2").alias("combined_values_2"))
dfOutput.show()
【问题讨论】:
标签: scala apache-spark apache-spark-sql