如何在PySpark DataFrame中对字符串数组类型列进行过滤?
解决PySpark数组列的过滤问题
针对你这个需求,其实有几种简洁的实现方式,我给你详细说下:
方法1:直接使用Spark数组字面量进行相等判断
PySpark支持数组类型的直接相等比较,但你不能直接用Python的列表['a','b']去和Spark的数组列比较,得用Spark内置的array()函数来构造对应的数组表达式,结合lit()来包装字符串元素。
完整代码示例:
from pyspark.sql import SparkSession from pyspark.sql.functions import array, lit # 初始化SparkSession(如果还没初始化的话) spark = SparkSession.builder.appName("array-filter").getOrCreate() # 你的原始DataFrame df = spark.createDataFrame([ ("u1", ['a', 'b']), ("u2", ['c', 'b']), ("u3", ['a', 'b']), ], ['user_id', 'features']) # 过滤features等于['a','b']的行 filtered_df = df.filter(df.features == array(lit('a'), lit('b'))) # 查看结果 filtered_df.show(truncate=False)
方法2:使用array_equal()函数(Spark 2.4+版本支持)
如果你使用的Spark版本在2.4及以上,可以用更语义化的array_equal()函数,它专门用来判断两个数组是否完全相等(元素顺序和值都一致):
from pyspark.sql.functions import array_equal, array, lit filtered_df = df.filter(array_equal(df.features, array(lit('a'), lit('b')))) filtered_df.show(truncate=False)
执行后的预期输出
两种方法都会得到你想要的结果:
+-------+--------+ |user_id|features| +-------+--------+ |u1 |[a, b] | |u3 |[a, b] | +-------+--------+
补充说明
为什么不能直接用df.filter(df.features == ['a','b'])?因为Python的列表对象和Spark的ArrayType数据类型不兼容,Spark无法识别Python列表作为SQL表达式的一部分,必须用Spark提供的函数来构造对应的数组表达式才行。
内容的提问来源于stack exchange,提问作者n0obcoder
相关产品推荐
相关产品推荐

