如何用Python循环实现DataFrame按id分组的子图绘制?
问题:简化多分组子图的重复绘制代码
现有如下Pandas DataFrame,需要为每个id分组绘制cycle与Salary的关系子图,当前通过手动重复编写4次subplot代码实现,希望通过迭代逻辑减少重复代码。
原始DataFrame代码
# Load the required libraries import pandas as pd import matplotlib.pyplot as plt # Create dataset data = {'id': [1, 1, 1, 1, 1, 1,1, 1, 1, 1, 1, 1, 2, 2, 2, 2, 2, 2, 2, 3, 3, 3, 3, 3, 3, 3,3, 4, 4, 4, 4, 4,4,], 'cycle': [0.0, 0.2,0.4, 0.6, 0.8, 1,1.2,1.4,1.6,1.8,2.0,2.2, 0.0, 0.2,0.4, 0.6,0.8,1.0,1.2, 0.0, 0.2,0.4, 0.6, 0.8,1.0,1.2,1.4, 0.0, 0.2,0.4, 0.6, 0.8,1.0,], 'Salary': [6, 7, 7, 7,8,9,10,11,12,13,14,15, 3, 4, 4, 4,4,5,6, 2, 8,9,10,11,12,13,14, 1, 8,9,10,11,12,], 'Children': ['Yes', 'No', 'Yes', 'Yes', 'Yes', 'Yes', 'No','No', 'Yes', 'Yes', 'Yes', 'No', 'Yes', 'Yes', 'Yes', 'No', 'Yes', 'Yes', 'Yes', 'Yes', 'No','Yes', 'Yes', 'No','No', 'Yes','Yes', 'Yes', 'Yes', 'No','Yes', 'Yes','Yes',], 'Days': [141, 123, 128, 66, 66, 120, 141, 52,96, 120, 141, 52, 141, 96, 120,120, 141, 52,96, 141, 15,123, 128, 66, 120, 141, 141, 141, 141,123, 128, 66,67,], } # Convert to dataframe df = pd.DataFrame(data)
简化实现方法
可以通过迭代分组数据的方式,只编写一次绘图逻辑即可完成所有子图的绘制,以下提供两种常用实现方式:
方式一:直接遍历groupby分组对象
这种方法自动处理所有id分组,无需手动指定数量,扩展性更强:
plt_fig_verify = plt.figure(figsize=(10,8)) # 按id分组 grouped_data = df.groupby('id') # 获取分组数量,用于设置子图布局 group_count = len(grouped_data) # 迭代每个分组 for plot_idx, (id_value, sub_df) in enumerate(grouped_data, start=1): # 创建子图:行数为分组数,列数为1,当前子图索引为plot_idx plt.subplot(group_count, 1, plot_idx) # 绘制cycle与Salary的曲线 plt.plot(sub_df['cycle'], sub_df['Salary'], 'b', linewidth=1, label=f'id{id_value}') # 设置标签与图例 plt.xlabel('cycle') plt.ylabel('Salary') plt.legend() # 自动调整子图间距,避免标签重叠 plt.tight_layout() plt.show()
方式二:遍历唯一id列表
这种方法逻辑直观,适合需要对id做额外预处理的场景:
plt_fig_verify = plt.figure(figsize=(10,8)) # 获取所有唯一的id值 unique_ids = df['id'].unique() id_count = len(unique_ids) # 迭代每个id for plot_idx, id_value in enumerate(unique_ids, start=1): # 筛选当前id对应的数据集 sub_df = df[df['id'] == id_value] # 创建子图 plt.subplot(id_count, 1, plot_idx) # 绘制曲线 plt.plot(sub_df['cycle'], sub_df['Salary'], 'b', linewidth=1, label=f'id{id_value}') plt.xlabel('cycle') plt.ylabel('Salary') plt.legend() plt.tight_layout() plt.show()
内容的提问来源于stack exchange,提问作者NN_Developer
相关产品推荐
相关产品推荐

