ENH: pass list-like as categories to plot (#1173)

This commit is contained in:
Martin Fleischmann authored and GitHub committed 2020-06-22 22:00:41 +02:00
1 parent bbf06a753c
commit caefd7562a
2 files changed
+53 -2

No files matched your search

+26 -2
View File
@@ -448,6 +448,7 @@ def plot_dataframe(
markersize=None,
figsize=None,
legend_kwds=None,
categories=None,
classification_kwds=None,
missing_kwds=None,
aspect="auto",
@@ -514,6 +515,8 @@ def plot_dataframe(
legend_kwds : dict (default None)
Keyword arguments to pass to matplotlib.pyplot.legend() or
matplotlib.pyplot.colorbar().
categories : list-like
Ordered list-like object of categories to be used for categorical plot.
classification_kwds : dict (default None)
Keyword arguments to pass to mapclassify
missing_kwds : dict (default None)
@@ -625,14 +628,35 @@ def plot_dataframe(
nan_idx = pd.isna(values)
if categories:
categorical = True
# Define `values` as a Series
if categorical:
if cmap is None:
cmap = "tab10"
categories = list(set(values[~nan_idx]))
categories.sort()
if categories is None:
categories = list(set(values[~nan_idx]))
categories.sort()
else:
if not pd.api.types.is_list_like(categories):
raise ValueError(
"Categories must be a list-like object. Objects that are "
"considered list-like are for example Python lists, tuples, sets, "
"NumPy arrays, and Pandas Series."
)
missing = [x for x in values[~nan_idx] if x not in categories]
if missing:
raise ValueError(
"Column contains values not listed in categories. "
"Missing categories: {}.".format(missing)
)
valuemap = dict((k, v) for (v, k) in enumerate(categories))
values = np.array([valuemap[k] for k in values[~nan_idx]])
vmin = 0 if vmin is None else vmin
vmax = max(valuemap.values()) if vmax is None else vmax
if scheme is not None:
if classification_kwds is None:
+27
View File
@@ -262,6 +262,33 @@ class TestPointPlotting:
with pytest.raises(TypeError): # no list allowed for alpha
ax = self.df2.plot(alpha=[0.7, 0.2])
def test_categorical_list(self):
self.df["cats"] = ["cat1", "cat2"] * 5
self.df["nums"] = [1, 2] * 5
self.df["singlecat"] = ["cat2"] * 10
ax1 = self.df.plot("cats", legend=True)
ax2 = self.df.plot("singlecat", categories=["cat1", "cat2"], legend=True)
ax3 = self.df.plot("nums", categories=[1, 2], legend=True)
point_colors1 = ax1.collections[0].get_facecolors()
point_colors2 = ax2.collections[0].get_facecolors()
point_colors3 = ax3.collections[0].get_facecolors()
np.testing.assert_array_equal(
point_colors1[1], point_colors2[0], point_colors3[0]
)
legend1 = [x.get_markerfacecolor() for x in ax1.get_legend().get_lines()]
legend2 = [x.get_markerfacecolor() for x in ax2.get_legend().get_lines()]
legend3 = [x.get_markerfacecolor() for x in ax3.get_legend().get_lines()]
np.testing.assert_array_equal(legend1, legend2)
np.testing.assert_array_equal(legend1, legend3)
with pytest.raises(ValueError, match="Categories must be a list-like object."):
self.df.plot(column="cats", categories="non_list")
with pytest.raises(
ValueError, match="Column contains values not listed in categories."
):
self.df.plot(column="cats", categories=["cat1"])
def test_misssing(self):
self.df.loc[0, "values"] = np.nan
ax = self.df.plot("values")