Plotting by grouped data using Matplotlib

Viewed 161

I'm using Matplotlib and Pandas to plot x by y, grouped by z. So I have the following:

x = df['ColumnA']
y = df['ColumnB']
fig, ax = plt.subplots(figsize=(20, 10))
for key, grp in df.groupby(['ColumnC']):
    plt.plot(grp['ColumnA'], grp['ColumnB'].rolling(window=30).mean(), label=key)

I also want to highlight 2 specific values from the total amount of values that will be plotted:

ax.legend(('Value1', 'Value2'))
plt.show()

This works fine. I just have the 2 values in my legend, but all values are actually plotted. What I actually want, is to be able to specify the colors for the 2 Values above. i.e. red and blue and have all the other values from Column C show on the plot as one color. The objective is to highlight how Value 1 & 2 are performing compared to everything else.

1 Answers

First, change the colours of the lines of interest.

lines_to_highlight = {
    'Value1': 'red',
    'Value2': 'blue'
}
DEFAULT_COLOR = 'gray'

legend_entries = []  # Lines to show in legend

for line in ax.lines:
    if line.get_label() in lines_to_highlight:
        line.set_color(lines_to_highlight[line.get_label()])

        legend_entries.append(line)
    else:
        line.set_color(DEFAULT_COLOR)

Second, create your legend.

ax.legend(
    legend_entries, 
    [entry.get_label() for entry in legend_entries]
)

Notes:

  • ax.legend(('Value1', 'Value2')) doesn't do what you expect. It simply resets the labels for the first two lines you plotted. It doesn't restrict the legend to lines you created with those labels. (The matplotlib docs themselves say that this mistake is easy to make.)
  • You must call ax.legend(...) after setting the line colours. Otherwise, the colours in the legend might not match the ones in the plot.

Example

import matplotlib.pyplot as plt

ax = plt.subplot(111)
ax.plot([1, 1, 1], label='one')
ax.plot([2, 2, 2], label='two')
ax.plot([3, 3, 3], label='three')

lines_to_highlight = {
    'one': 'red',
    'three': 'blue'
}
DEFAULT_COLOR = 'gray'

legend_entries = []  # Lines to show in legend

for line in ax.lines:
    if line.get_label() in lines_to_highlight:
        line.set_color(lines_to_highlight[line.get_label()])

        legend_entries.append(line)
    else:
        line.set_color(DEFAULT_COLOR)

ax.legend(
    legend_entries, [entry.get_label() for entry in legend_entries]
)
plt.show()

Minimal example

Related