diff --git a/geopandas/plotting.py b/geopandas/plotting.py index 9076113..169fe39 100644 --- a/geopandas/plotting.py +++ b/geopandas/plotting.py @@ -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: diff --git a/geopandas/tests/test_plotting.py b/geopandas/tests/test_plotting.py index 2baf53f..2dec35e 100644 --- a/geopandas/tests/test_plotting.py +++ b/geopandas/tests/test_plotting.py @@ -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")