PySpark DataFrame按id1分组生成堆叠柱状图求助
Databricks PySpark按id1分组生成堆叠柱状图解决方案
问题背景
在Databricks的PySpark环境中,需处理1.3亿行的长格式DataFrame(含1200个唯一id1值),为每个id1分组生成独立堆叠柱状图,y轴对应fields列的类别。尝试转Pandas DataFrame时出现**OOM(内存不足)**错误,参考Stack Overflow相关问题后仍未写出可运行代码。
数据示例
data = [{'id1': '1150', 'org':'south_org', 'id2':300, 'fields':'firstname', 'value':'complete'}, {'id1': '1150', 'org':'south_org', 'id2':300, 'fields':'lastname', 'value':'complete'}, {'id1': '1150', 'org':'south_org', 'id2':300, 'fields':'streetaddress', 'value':'missing'}, {'id1': '1150', 'org':'south_org', 'id2':300, 'fields':'state', 'value':'missing'}, {'id1': '1150', 'org':'south_org', 'id2':300, 'fields':'zip', 'value':'complete'}, {'id1': '3000', 'org':'north_org', 'id2':310, 'fields':'firstname', 'value':'complete'}, {'id1': '3000', 'org':'north_org', 'id2':310, 'fields':'lastname', 'value':'complete'}, {'id1': '3000', 'org':'north_org', 'id2':310, 'fields':'streetaddress', 'value':'error'}, {'id1': '3000', 'org':'north_org', 'id2':310, 'fields':'state', 'value':'complete'}, {'id1': '3000', 'org':'north_org', 'id2':310, 'fields':'zip', 'value':'complete'}, {'id1': '1110', 'org':'west_org', 'id2':315, 'fields':'firstname', 'value':'complete'}, {'id1': '1110', 'org':'west_org', 'id2':315, 'fields':'lastname', 'value':'complete'}, {'id1': '1110', 'org':'west_org', 'id2':315, 'fields':'streetaddress', 'value':'complete'}, {'id1': '1110', 'org':'west_org', 'id2':315, 'fields':'state', 'value':'complete'}, {'id1': '1110', 'org':'west_org', 'id2':315, 'fields':'zip', 'value':'complete'} ] # 转换为PySpark DataFrame df = spark.createDataFrame(data)
之前尝试的报错代码及问题
代码1
grouped = data.groupby('id1') for i, (groupname, group) in enumerate(grouped): axes[i].plot(group.id1, group.values, kind='stacked') axes[i].set_title(groupname)
问题:PySpark的groupby返回GroupedData对象,无法像Pandas直接迭代分组数据;绘图参数错误,未指定堆叠柱状图的正确格式。
代码2
import matplotlib.pyplot as plt import seaborn as sb for id1, data in wide_df: plt.figure(figsize=(20,5)) sns.barplot(data=data, x='fields',y='status') plt.tight_layout() plt.show()
问题:wide_df未正确生成;sns.barplot默认不支持堆叠,且参数y='status'在原数据中不存在,应为处理后的计数列。
可行解决方案
利用PySpark分布式计算先做数据聚合,再逐个提取单个id1的小批量数据转Pandas绘图,避免内存溢出:
1. 导入依赖库
from pyspark.sql import functions as F import matplotlib.pyplot as plt
2. Spark端数据预处理(核心:避免全量加载内存)
# 统计每个id1下,各字段(fields)对应状态(value)的数量 agg_df = df.groupBy("id1", "fields", "value").agg(F.count("*").alias("count")) # 转宽表:每个字段为一行,不同状态为列,值为对应计数 pivot_df = agg_df.groupBy("id1", "fields").pivot("value").agg(F.sum("count")).fillna(0)
3. 遍历每个id1生成堆叠柱状图
# 获取所有唯一id1值 id_list = [row.id1 for row in df.select("id1").distinct().collect()] # 逐个处理每个id1 for id_val in id_list: # 过滤当前id1的数据,转换为小Pandas DataFrame(仅包含当前id1的少量行) single_id_df = pivot_df.filter(F.col("id1") == id_val).drop("id1").toPandas() # 设置字段名为索引,方便绘图 single_id_df.set_index("fields", inplace=True) # 绘制堆叠柱状图 plt.figure(figsize=(10, 6)) single_id_df.plot(kind="bar", stacked=True, ax=plt.gca()) plt.title(f"id1: {id_val} - 字段状态分布") plt.xlabel("字段名称") plt.ylabel("数量") plt.xticks(rotation=45) plt.legend(title="状态") plt.tight_layout() plt.show()
关键说明
- 所有聚合、透视操作在Spark分布式集群完成,不会将1.3亿行全量数据加载到单节点内存;
- 每个id1对应的Pandas DataFrame仅包含该id1的字段数据(示例中为5行),内存占用极低,彻底避免OOM;
- 通过Pandas的
plot(kind='bar', stacked=True)实现堆叠柱状图,符合需求的可视化效果。
内容的提问来源于stack exchange,提问作者budding pro
相关产品推荐
相关产品推荐

