Python How to plot "grouped by" scatter on (mean of column vs percentile)

Viewed 424

I have the following dataframe df_male:

enter image description here

And I used groupby() to see mean value of gagne_sum_t column on each risk_percentile, df_male.groupby(["risk_percentile","race"]).aggregate(np.mean):

enter image description here

I want to scatterplot this gagne_sum_t vs risk_percentile grouped by race, for something like: enter image description here

With this legend for the plot: enter image description here

However, I am not too sure how to proceed from here... How do I use groupby() again from here?

1 Answers

I don't think you need to do group_by again. You can start plotting from the grouped dataframe. You already have this:

plot_df = df_male.groupby(["risk_percentile","race"]).aggregate(np.mean)

If you have pyplot (typically imported as plt) then you can plot straight from that dataframe:

import matplotlib.pyplot as plt

fig, ax = plt.subplots()
ax.scatter(x=plot_df.gagne_sum_t, y=plot_df.index.names[0], c=plot_df.race)
ax.legend() 
plt.show()
Related