【问题标题】:get first N elements from dataframe ArrayType column in pyspark从pyspark中的数据框ArrayType列中获取前N个元素
【发布时间】:2019-03-29 06:10:00
【问题描述】:

我有一个 spark 数据框,其行为 -

1   |   [a, b, c]
2   |   [d, e, f]
3   |   [g, h, i]

现在我只想保留数组列中的前 2 个元素。

1   |   [a, b]
2   |   [d, e]
3   |   [g, h]

如何实现?

注意 - 请记住,我在这里提取的不是单个数组元素,而是可能包含多个元素的数组的一部分。

【问题讨论】:

  • 我已经看到了那个答案,但这不是我想要的。我不想要数组中的单个项目,而是寻找前 N 个元素。
  • @pault 有趣的是,链接解决方案似乎不适用于 Spark 2.3.1(引发异常)。有什么想法吗?
  • @pault 谜团解开了!一个新用户决定在链接答案中更改 OP 的代码,使其错误(已恢复)...
  • stackoverflow.com/questions/47585279/… 不相信我们需要为此创建一个临时视图

标签: apache-spark pyspark apache-spark-sql


【解决方案1】:

以下是使用 API 函数的方法。

假设您的 DataFrame 如下:

df.show()
#+---+---------+
#| id|  letters|
#+---+---------+
#|  1|[a, b, c]|
#|  2|[d, e, f]|
#|  3|[g, h, i]|
#+---+---------+

df.printSchema()
#root
# |-- id: long (nullable = true)
# |-- letters: array (nullable = true)
# |    |-- element: string (containsNull = true)

您可以使用方括号按索引访问letters 列中的元素,并将其包装在对pyspark.sql.functions.array() 的调用中以创建新的ArrayType 列。

import pyspark.sql.functions as f

df.withColumn("first_two", f.array([f.col("letters")[0], f.col("letters")[1]])).show()
#+---+---------+---------+
#| id|  letters|first_two|
#+---+---------+---------+
#|  1|[a, b, c]|   [a, b]|
#|  2|[d, e, f]|   [d, e]|
#|  3|[g, h, i]|   [g, h]|
#+---+---------+---------+

或者,如果要列出的索引过多,可以使用列表推导:

df.withColumn("first_two", f.array([f.col("letters")[i] for i in range(2)])).show()
#+---+---------+---------+
#| id|  letters|first_two|
#+---+---------+---------+
#|  1|[a, b, c]|   [a, b]|
#|  2|[d, e, f]|   [d, e]|
#|  3|[g, h, i]|   [g, h]|
#+---+---------+---------+

对于 pyspark 2.4+ 版本,您还可以使用 pyspark.sql.functions.slice():

df.withColumn("first_two",f.slice("letters",start=1,length=2)).show()
#+---+---------+---------+
#| id|  letters|first_two|
#+---+---------+---------+
#|  1|[a, b, c]|   [a, b]|
#|  2|[d, e, f]|   [d, e]|
#|  3|[g, h, i]|   [g, h]|
#+---+---------+---------+

slice 对于大型数组可能有更好的性能(注意起始索引是 1,而不是 0)

【讨论】:

  • 该死...正如我所说,我已经生锈了-甚至不记得pyspark.sql.functions.array存在 ... :(
  • 这对我来说不适用于类似的问题。我收到以下错误:“无法从概率#6225 中提取值:需要结构类型但得到了 struct,values:array>;”跨度>
  • @LePuppyle 我猜你有一个 VectorUDT 而不是一个数组。首先,你需要一个udf - 试试this post
  • AnalysisException: "Field name should be String Literal, but it's 0;"
【解决方案2】:

要么我的 pyspark 技能已经生疏(我承认我现在已经不再磨练它们了),要么这确实是一个棘手的问题......我设法做到这一点的唯一方法是使用 SQL 语句:

spark.version
#  u'2.3.1'

# dummy data:

from pyspark.sql import Row
x = [Row(col1="xx", col2="yy", col3="zz", col4=[123,234, 456])]
rdd = sc.parallelize(x)
df = spark.createDataFrame(rdd)
df.show()
# result:
+----+----+----+---------------+
|col1|col2|col3|           col4|
+----+----+----+---------------+
|  xx|  yy|  zz|[123, 234, 456]|
+----+----+----+---------------+

df.createOrReplaceTempView("df")
df2 = spark.sql("SELECT col1, col2, col3, (col4[0], col4[1]) as col5 FROM df")
df2.show()
# result:
+----+----+----+----------+ 
|col1|col2|col3|      col5|
+----+----+----+----------+ 
|  xx|  yy|  zz|[123, 234]|
+----+----+----+----------+

对于未来的问题,最好遵循How to make good reproducible Apache Spark Dataframe examples 上的建议指南。

【讨论】:

    猜你喜欢
    • 1970-01-01
    • 2016-10-14
    • 2020-05-28
    • 1970-01-01
    • 2021-12-22
    • 2014-09-17
    • 1970-01-01
    • 1970-01-01
    • 2016-04-25
    相关资源
    最近更新 更多