You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.22 00:36:28