如何使用Pandas按model与scheduler两列分组绘制mae列的多组条形图?
Hey there! The issue with your current code is that groupby().plot.bar() creates a separate subplot for each (model, scheduler) pair instead of grouping bars by model as you want. Let's fix this with two straightforward approaches:
Approach 1: Use Seaborn (Simplest for Grouped Plots)
Seaborn's catplot is built exactly for this kind of categorical grouped visualization. Here's how to use it:
First, let's make sure we have your DataFrame set up correctly (I'll recreate it from your table):
import pandas as pd import seaborn as sns import matplotlib.pyplot as plt # Recreate your DataFrame data = [ ["ecaresnet50t", "warm", 4.518], ["ecaresnet50t", "cosine", 4.46], ["ecaresnet50t", "constant", 4.972], ["resnest50d", "warm", 4.056], ["resnest50d", "cosine", 4.1], ["resnest50d", "constant", 5.072], ["resnetrs50", "warm", 4.164], ["resnetrs50", "cosine", 4.154], ["resnetrs50", "constant", 4.644], ["seresnet50", "warm", 4.202] ] df = pd.DataFrame(data, columns=["model", "scheduler", "mae"])
Now plot the grouped bar chart:
# Create the grouped bar plot sns.catplot( x="model", y="mae", hue="scheduler", kind="bar", data=df, palette="viridis" ) # Customize the plot for readability plt.title("MAE Scores by Model and Scheduler") plt.xlabel("Model") plt.ylabel("MAE Score") plt.xticks(rotation=45) # Rotate model names to avoid overlap plt.tight_layout() # Prevent label cutoff plt.show()
This will generate a single plot where each model has 3 distinct bars (one for each scheduler), colored differently for easy comparison.
Approach 2: Use Pandas Pivot + Plot
If you prefer sticking with pandas, reshape your DataFrame into a wide format first using pivot_table, then plot the bars:
# Reshape the DataFrame to wide format (model as rows, scheduler as columns) pivot_df = df.pivot_table(index="model", columns="scheduler", values="mae") # Plot the grouped bars pivot_df.plot(kind="bar", figsize=(10,6), palette="viridis") # Customize the plot plt.title("MAE Scores by Model and Scheduler") plt.xlabel("Model") plt.ylabel("MAE Score") plt.xticks(rotation=45) plt.legend(title="Scheduler") plt.tight_layout() plt.show()
This method organizes your data so each model is a row, each scheduler is a column, then plots each column as a bar group for the corresponding model.
Both approaches will give you the grouped bar chart you're looking for—pick whichever fits your workflow better!
内容的提问来源于stack exchange,提问作者Talha Anwar

