PySpark中如何获取collect_list()聚合后的嵌套列表值?
获取PySpark collect_list聚合后的嵌套列表值
方法1:直接从Row对象提取
你的代码执行collect()后返回的是包含单个Row对象的列表,只需通过索引定位Row,再通过聚合列的默认名称取值:
from pyspark.sql.functions import collect_list spark = SparkSession.builder.appName('example').getOrCreate() data = [{'users': '1', 'songs': 23}, {'users': '1', 'songs': 28}, {'users': '2', 'songs': 43}, {'users': '2', 'songs': 63}, {'users': '3', 'songs': 78}, {'users': '3', 'songs': 33}] dataframe = spark.createDataFrame(data) songs_mean = dataframe.groupBy('users').agg({'songs':'mean'}).agg(collect_list('avg(songs)')).collect() # 提取目标列表 target_list = songs_mean[0]['collect_list(avg(songs))'] print(target_list) # 输出: [55.5, 25.5, 53.0]
方法2:给聚合列起别名(可读性更强)
为避免使用冗长的默认列名,可给collect_list的结果指定别名:
from pyspark.sql.functions import collect_list # 为聚合列设置别名 result_df = dataframe.groupBy('users').agg({'songs':'mean'}).agg(collect_list('avg(songs)').alias('mean_songs_list')) # 提取列表 target_list = result_df.collect()[0]['mean_songs_list']
方法3:用first()替代collect()(更高效)
由于你的聚合结果仅一行数据,使用first()直接获取唯一Row对象,比collect()更高效:
target_list = dataframe.groupBy('users').agg({'songs':'mean'}).agg(collect_list('avg(songs)').alias('mean_songs_list')).first()['mean_songs_list']
原理说明
collect()会将DataFrame所有行以Row对象列表的形式返回,这里聚合结果仅一行,因此取列表索引[0]即可得到目标Row对象。- Row对象支持通过列名字符串索引访问对应值,从而直接提取嵌套的列表。
内容的提问来源于stack exchange,提问作者Simone Di Claudio
相关产品推荐
相关产品推荐

