如何在PySpark中将数据转换为ArrayType格式并保存为CSV?
PySpark实现用户卡片数据按ID聚合为数组格式并保存CSV
问题需求
现有如下格式的用户卡片数据:
cust_id | card_num | balance | payment | due | card_type |
|---|---|---|---|---|---|
| c1 | 1234 | 567 | 344 | 33 | A |
| c1 | 2345 | 57 | 44 | 3 | B |
| c2 | 123 | 561 | 34 | 39 | A |
| c3 | 345 | 517 | 914 | 23 | C |
| c3 | 127 | 56 | 34 | 32 | B |
| c3 | 347 | 67 | 344 | 332 | B |
需要将其转换为按cust_id分组,其余字段聚合为数组的格式,目标输出如下:
cust_id | card_num | balance | payment | due | card_type |
|---|---|---|---|---|---|
| c1 | [1234, 2345] | [567, 57] | [344, 44] | [33, 3] | [A, B] |
| c2 | [123] | [561] | [34] | [39] | [A] |
| c3 | [345, 127, 347] | [517, 56, 67] | [914, 34, 344] | [23, 32, 332] | [C, B, B] |
通用PySpark实现代码
以下是无需硬编码非分组字段的通用实现:
from pyspark.sql import SparkSession from pyspark.sql.functions import collect_list # 1. 初始化SparkSession spark = SparkSession.builder \ .appName("CustCardAggregation") \ .getOrCreate() # 2. 读取输入数据(示例为CSV格式,可根据实际数据源调整) input_df = spark.read.csv( path="path/to/your/input.csv", header=True, inferSchema=True ) # 3. 定义分组键与待聚合列 group_key = "cust_id" agg_columns = [col for col in input_df.columns if col != group_key] # 4. 分组聚合:将每个非分组列转为数组 agg_exprs = [collect_list(col).alias(col) for col in agg_columns] result_df = input_df.groupBy(group_key).agg(*agg_exprs) # 5. 保存结果为CSV result_df.write.csv( path="path/to/your/output.csv", header=True, mode="overwrite" # 可选:覆盖已有输出目录,可替换为append/ignore ) # 关闭SparkSession spark.stop()
代码说明
- SparkSession初始化:创建Spark应用入口,可按需添加集群配置、资源分配等参数。
- 数据读取:示例以CSV为例,
inferSchema=True自动推断字段类型;若需精确控制字段类型,可手动定义schema参数。 - 通用聚合逻辑:通过列表推导式自动识别所有非分组字段,生成
collect_list聚合表达式,避免逐个字段编写,保证代码通用性。 - 结果保存:
header=True保留表头,mode参数可根据需求调整为追加、忽略等模式。
内容的提问来源于stack exchange,提问作者suraj jadhav
相关产品推荐
相关产品推荐

