Seaborn relplot分面子图添加各分组移动平均线方法
seaborn分面图叠加分组移动平均线实现方案
问题概述
- 核心需求:在按
customer字段拆分的24个seaborn relplot分钟级时间序列子图中,为每个子图叠加对应客户分组的24小时移动平均线 - 移动平均计算规则:对
avg_price字段做窗口为24小时、最小有效观测数360的滚动均值,逻辑为df['avg_price'].rolling('24H', min_periods=360).mean() - 已尝试的无效操作:
- 在
for ax in g.axes.flat循环中直接传入全量数据集的滚动均值,用ax.plot()绘图 - 循环内调用
sns.lineplot()绘制全量滚动均值结果 - 已成功为DataFrame添加分组计算的MA列,但未找到正确的seaborn调用方式
- 在
- 现有基础绘图代码:
facet_kws={'sharey': False} g= sns.relplot( data=df, x=df.index, y="avg_price", col="customer", kind="line", col_wrap=4, color='k', alpha=0.75, facet_kws=facet_kws, ); g.set_axis_labels("Date", "Average price", fontsize=13); g.set_titles("{col_name}", size=14); g.set_xticklabels(rotation=60, fontsize=13); ax = plt.gca() ax.xaxis.set_major_locator(md.HourLocator(interval=24)) for ax in g.axes.flat: ax.axvspan(pd.Timestamp(start_test), pd.Timestamp(end_test), color='y', alpha=0.25, lw=0);
- 数据集结构:索引为
date_time时间字段,样例如下:
customer avg_price avg_price2 count1 count2 date_time 2022-06-11 00:00:00 Customer1 4.4656 1.25 36 11084 2022-06-11 00:00:00 Customer2 7.8873 0.92 10 22150 2022-06-11 00:00:00 Customer3 2.3016 1.37 1 2521 2022-06-11 00:00:00 Customer4 3.2421 1.05 221 98973 2022-06-11 00:00:00 Customer5 1.0050 0.94 2 410 ... ... ... ... ... ... ... 2022-06-21 10:00:00 Customer1 4.9450 1.99 340 118000 2022-06-21 10:00:00 Customer2 4.0643 2.06 268 20850 2022-06-21 10:00:00 Customer3 3.7034 1.00 25 5100 2022-06-21 10:00:00 Customer4 5.0367 2.64 2098 118251 2022-06-21 10:00:00 Customer5 2.7429 1.57 50 11900
失效原因
之前的方法核心问题是循环绘图时传入的是全量数据集的MA值,没有按当前子图对应的customer做数据过滤,导致所有客户的时间点数据混在一起,要么线条完全错乱,要么因重复数据过多无法正常渲染。另外原代码中x轴时间定位器的设置仅作用于最后一个生成的子图,没有覆盖全部分面。
可行实现方案
方案1:预计算分组MA列,循环子图时过滤对应数据绘图(最简便)
- 首先确保时间索引升序,计算对齐原数据的分组MA列:
# 时间序列滚动计算必须先按时间排序 df = df.sort_index() # 按客户分组计算24H滚动均值,结果和原数据行对齐 df['MA_24H'] = df.groupby('customer')['avg_price'].transform( lambda x: x.rolling('24H', min_periods=360).mean() )
- 修改绘图代码,遍历子图时匹配对应客户数据绘制MA:
import seaborn as sns import pandas as pd import matplotlib.dates as md import matplotlib.pyplot as plt facet_kws={'sharey': False} g= sns.relplot( data=df, x=df.index, y="avg_price", col="customer", kind="line", col_wrap=4, color='k', alpha=0.75, facet_kws=facet_kws, ); g.set_axis_labels("Date", "Average price", fontsize=13); g.set_titles("{col_name}", size=14); g.set_xticklabels(rotation=60, fontsize=13); for ax in g.axes.flat: # 从子图标题获取当前分面对应的客户名 current_customer = ax.get_title() # 过滤当前客户的子集数据 customer_subset = df[df['customer'] == current_customer] # 绘制对应MA线 ax.plot(customer_subset.index, customer_subset['MA_24H'], color='orange', label='24H MA', lw=1.5) # 保留原有黄色区间高亮逻辑 ax.axvspan(pd.Timestamp(start_test), pd.Timestamp(end_test), color='y', alpha=0.25, lw=0) # 统一设置x轴时间刻度间隔 ax.xaxis.set_major_locator(md.HourLocator(interval=24)) # 按需添加图例 ax.legend(fontsize=10) plt.tight_layout() plt.show()
方案2:长表转换后用seaborn原生映射绘图(无需手动循环)
将原始价格和MA值转换为长表格式,用hue参数区分线条类型,seaborn会自动完成分面匹配:
# 第一步同方案1,先计算分组MA列 df = df.sort_index() df['MA_24H'] = df.groupby('customer')['avg_price'].transform( lambda x: x.rolling('24H', min_periods=360).mean() ) # 转换为长表:合并原始价格和MA值到同一列,新增字段区分数据类型 df_long = df.melt( id_vars=['customer', 'avg_price2', 'count1', 'count2'], value_vars=['avg_price', 'MA_24H'], var_name='data_type', value_name='price', ignore_index=False # 保留原时间索引 ) # 定义颜色映射 palette = {'avg_price':'black', 'MA_24H':'orange'} # 直接绘图,hue自动区分原始线和MA线 g = sns.relplot( data=df_long, x=df_long.index, y='price', col='customer', hue='data_type', kind='line', col_wrap=4, palette=palette, facet_kws={'sharey':False}, alpha=0.75, linewidth=1 ) # 格式调整 g.set_axis_labels("Date", "Average price", fontsize=13) g.set_titles("{col_name}", size=14) g.set_xticklabels(rotation=60, fontsize=13) # 调整MA线宽度 for line in g.axes.flat[0].lines: if line.get_label() == 'MA_24H': line.set_linewidth(1.5) # 添加高亮区间和时间刻度 for ax in g.axes.flat: ax.axvspan(pd.Timestamp(start_test), pd.Timestamp(end_test), color='y', alpha=0.25, lw=0) ax.xaxis.set_major_locator(md.HourLocator(interval=24)) plt.tight_layout() plt.show()
注意事项:
- 时间窗口滚动计算前必须确认索引为datetime类型且升序排列,否则计算结果会出错
- 若MA线前期为空值是正常现象,是
min_periods=360参数要求的最小观测数未满足导致的- 分面图的全局格式设置必须遍历所有子图生效,仅对
plt.gca()设置只会作用于最后一个生成的子图
内容的提问来源于stack exchange,提问作者Geosphere
相关产品推荐
相关产品推荐

