diff --git a/CHANGELOG.md b/CHANGELOG.md index 4f4cc3ac0e..bc55cdf0d7 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -4,6 +4,9 @@ This project adheres to [Semantic Versioning](http://semver.org/). ## Unreleased +### Added +- Support `marginal_x`/`marginal_y="heatmap"` in `density_heatmap`, drawing a single-row/column heatmap strip in the margin colored by the same `z`/`histfunc` aggregate as the main plot and sharing its color scale; `text_auto` now also applies to the marginal heatmap strip [[#5706](https://github.com/plotly/plotly.py/issues/5706)] + ### Fixed - Fix `hex_to_rgb` parsing of 3-digit shorthand hexadecimal colors such as `#FFF` [[#5662](https://github.com/plotly/plotly.py/pull/5662)], with thanks to @genrichez for the contribution! - Add `` to the `to_html()` template to comply with modern web standards [[#5693](https://github.com/plotly/plotly.py/pull/5693)], with thanks to @mishrakushal for the contribution! diff --git a/doc/python/marginal-plots.md b/doc/python/marginal-plots.md index 49728403f7..bd1ff69e74 100644 --- a/doc/python/marginal-plots.md +++ b/doc/python/marginal-plots.md @@ -39,7 +39,7 @@ Marginal distribution plots are small subplots above or to the right of a main p ### Scatter Plot Marginals -The `marginal_x` and `marginal_y` arguments accept one of `"histogram"`, `"rug"`, `"box"`, or `"violin"` (see also how to create [histograms](/python/histograms/), [box plots](/python/box-plots/) and [violin plots](/python/violin-plots/) as the main figure). +The `marginal_x` and `marginal_y` arguments accept one of `"histogram"`, `"rug"`, `"box"`, or `"violin"` (see also how to create [histograms](/python/histograms/), [box plots](/python/box-plots/) and [violin plots](/python/violin-plots/) as the main figure), plus `"heatmap"` for `density_heatmap` (see below). Marginal plots are linked to the main plot: try zooming or panning on the main plot. @@ -59,6 +59,27 @@ fig = px.density_heatmap(df, x="sepal_length", y="sepal_width", marginal_x="box" fig.show() ``` +### Marginal Heatmaps on Density Heatmaps + +`marginal_x` and `marginal_y` also accept `"heatmap"` for [`density_heatmap`](/python/2D-Histogram/). This draws a single-row or single-column heatmap strip, colored by the same aggregate (`histfunc` of `z`, or count by default) as the main plot, and sharing its color scale. This is not supported for `density_contour`, since a contour plot's colorbar is discrete and not compatible with the continuous colorbar used by the marginal heatmap. + +```python +import plotly.express as px +df = px.data.tips() +fig = px.density_heatmap(df, x="total_bill", y="tip", marginal_x="heatmap", marginal_y="heatmap") +fig.show() +``` + +Set `text_auto=True` to display the aggregate value as text on both the main plot and the marginal heatmap strips, or pass a [d3-format](https://github.com/d3/d3-format) string such as `".2f"` to control the number of decimal places: + +```python +import plotly.express as px +df = px.data.tips() +fig = px.density_heatmap(df, x="total_bill", y="tip", z="size", histfunc="avg", + marginal_x="heatmap", marginal_y="heatmap", text_auto=".1f") +fig.show() +``` + ### Marginal Plots and Color Marginal plots respect the `color` argument as well, and are linked to the respective legend elements. Try clicking on the legend items. diff --git a/plotly/express/_chart_types.py b/plotly/express/_chart_types.py index 2b4b10184e..426dd76eca 100644 --- a/plotly/express/_chart_types.py +++ b/plotly/express/_chart_types.py @@ -214,6 +214,15 @@ def density_heatmap( histfunc=[ "The arguments to this function are the values of `z`.", ], + marginal_x=[ + "Also supports `'heatmap'`, showing a single-row heatmap colored by the aggregate value.", + ], + marginal_y=[ + "Also supports `'heatmap'`, showing a single-column heatmap colored by the aggregate value.", + ], + text_auto=[ + "Also applies to `marginal_x`/`marginal_y='heatmap'`, in which case the z values are always displayed.", + ], ), ) diff --git a/plotly/express/_core.py b/plotly/express/_core.py index 1b599cb43c..2f00be9de5 100644 --- a/plotly/express/_core.py +++ b/plotly/express/_core.py @@ -971,12 +971,55 @@ def make_trace_spec(args, constructor, attrs, trace_patch): ), marginal=letter, ) + elif args["marginal_" + letter] == "heatmap": + if constructor != go.Histogram2d: + raise ValueError( + "`marginal_x`/`marginal_y` value `'heatmap'` is only supported " + "for `density_heatmap`." + ) + other_letter = "y" if letter == "x" else "x" + heatmap_trace_patch = dict( + coloraxis="coloraxis1", histfunc=args.get("histfunc"), **axis_map + ) + # `nbinsx`/`nbinsy` are only a target bin count -- plotly.js's "nice + # number" bin-sizing can still round to more than one bin. Force + # exactly one bin by setting explicit bin edges covering the data. + other_col = args["data_frame"].get_column(args[other_letter]) + other_min = nw.to_py_scalar(other_col.min()) + other_max = nw.to_py_scalar(other_col.max()) + span = (other_max - other_min) or 1 + pad = span * 0.001 + other_bins = dict( + start=other_min - pad, end=other_max + pad, size=span + 2 * pad + ) + if letter == "x": + heatmap_trace_patch["xbingroup"] = "x" + heatmap_trace_patch["ybins"] = other_bins + else: + heatmap_trace_patch["ybingroup"] = "y" + heatmap_trace_patch["xbins"] = other_bins + if args.get("text_auto", False) is not False: + if args["text_auto"] is True: + heatmap_trace_patch["texttemplate"] = "%{z}" + else: + heatmap_trace_patch["texttemplate"] = ( + "%{z:" + args["text_auto"] + "}" + ) + trace_spec = TraceSpec( + constructor=go.Histogram2d, + attrs=[letter, other_letter, "z"], + trace_patch=heatmap_trace_patch, + marginal=letter, + ) else: raise ValueError( f"Invalid value '{args['marginal_' + letter]}' for `marginal_{letter}`. " - "Supported marginal plot types are: 'rug', 'box', 'violin', 'histogram'." + "Supported marginal plot types are: " + "'rug', 'box', 'violin', 'histogram', 'heatmap'." ) - if "color" in attrs or "color" not in args: + if trace_spec.constructor != go.Histogram2d and ( + "color" in attrs or "color" not in args + ): if "marker" not in trace_spec.trace_patch: trace_spec.trace_patch["marker"] = dict() first_default_color = args["color_continuous_scale"][0] diff --git a/tests/test_optional/test_px/test_marginals.py b/tests/test_optional/test_px/test_marginals.py index 40d3019f2d..7e3278b1c5 100644 --- a/tests/test_optional/test_px/test_marginals.py +++ b/tests/test_optional/test_px/test_marginals.py @@ -26,6 +26,106 @@ def test_single_marginals(backend, px_fn, marginal, orientation): assert len(fig.data) == 1 + (marginal is not None) +def test_marginal_heatmap_uses_z_and_histfunc(backend): + df = px.data.tips(return_type=backend) + # backend-independent reference for min/max, since e.g. pyarrow columns don't + # support .min()/.max() directly + pdf = px.data.tips() + + fig = px.density_heatmap( + df, + x="total_bill", + y="tip", + z="size", + histfunc="sum", + marginal_x="heatmap", + marginal_y="heatmap", + ) + assert len(fig.data) == 3 + marginal_x_trace, marginal_y_trace = fig.data[1], fig.data[2] + + assert marginal_x_trace.type == "histogram2d" + assert marginal_x_trace.coloraxis == "coloraxis" + assert marginal_x_trace.histfunc == "sum" + # a single bin covering the full y range, so the strip is exactly one row + assert marginal_x_trace.ybins.start <= pdf["tip"].min() + assert marginal_x_trace.ybins.end >= pdf["tip"].max() + assert marginal_x_trace.ybins.size >= pdf["tip"].max() - pdf["tip"].min() + + assert marginal_y_trace.type == "histogram2d" + assert marginal_y_trace.coloraxis == "coloraxis" + assert marginal_y_trace.histfunc == "sum" + # a single bin covering the full x range, so the strip is exactly one column + assert marginal_y_trace.xbins.start <= pdf["total_bill"].min() + assert marginal_y_trace.xbins.end >= pdf["total_bill"].max() + assert ( + marginal_y_trace.xbins.size >= pdf["total_bill"].max() - pdf["total_bill"].min() + ) + + assert fig.layout.coloraxis.colorbar.title.text == "sum of size" + + +def test_marginal_heatmap_without_z(backend): + df = px.data.tips(return_type=backend) + + fig = px.density_heatmap( + df, x="total_bill", y="tip", marginal_x="heatmap", marginal_y="heatmap" + ) + marginal_x_trace, marginal_y_trace = fig.data[1], fig.data[2] + + assert marginal_x_trace.type == "histogram2d" + assert marginal_x_trace.coloraxis == "coloraxis" + assert marginal_x_trace.histfunc is None + + assert marginal_y_trace.type == "histogram2d" + assert marginal_y_trace.coloraxis == "coloraxis" + assert marginal_y_trace.histfunc is None + + assert fig.layout.coloraxis.colorbar.title.text == "count" + + +@pytest.mark.parametrize("text_auto", [True, ".1f"]) +def test_marginal_heatmap_text_auto(backend, text_auto): + df = px.data.tips(return_type=backend) + + fig = px.density_heatmap( + df, + x="total_bill", + y="tip", + marginal_x="heatmap", + marginal_y="heatmap", + text_auto=text_auto, + ) + expected = "%{z}" if text_auto is True else "%{z:" + text_auto + "}" + for trace in fig.data: + assert trace.texttemplate == expected + + +def test_marginal_heatmap_no_text_auto(backend): + df = px.data.tips(return_type=backend) + + fig = px.density_heatmap( + df, x="total_bill", y="tip", marginal_x="heatmap", marginal_y="heatmap" + ) + for trace in fig.data: + assert trace.texttemplate is None + + +def test_marginal_heatmap_unsupported_chart_type_raises(): + with pytest.raises(ValueError, match="only supported for `density_heatmap`"): + px.scatter(x=[1, 2, 3], y=[2, 3, 4], marginal_x="heatmap") + with pytest.raises(ValueError, match="only supported for `density_heatmap`"): + px.scatter(x=[1, 2, 3], y=[2, 3, 4], marginal_y="heatmap") + with pytest.raises(ValueError, match="only supported for `density_heatmap`"): + px.histogram(x=[1, 2, 3], marginal="heatmap") + # density_contour's discrete-looking colorbar isn't compatible with the + # continuous marginal heatmap colorbar, so it's intentionally unsupported + with pytest.raises(ValueError, match="only supported for `density_heatmap`"): + px.density_contour(x=[1, 2, 3], y=[2, 3, 4], marginal_x="heatmap") + with pytest.raises(ValueError, match="only supported for `density_heatmap`"): + px.density_contour(x=[1, 2, 3], y=[2, 3, 4], marginal_y="heatmap") + + def test_unsupported_marginal_raises_clear_error(): # issue 4654 # An unsupported marginal type used to fail deep inside make_figure with a # cryptic "'NoneType' object has no attribute 'constructor'". It should