FazBrowse GitHub Viewer | Trending |
URL:
| Home
Tools: [Download Repo ZIP]   [Original HTTPS Page]

ENH: Added sharex/sharey string support to subplot_mosaic by nillohitroy · Pull Request #32437 · matplotlib/matplotlib · GitHub

Repository navigation

Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension .py  (3) .pyi  (1) All 2 file types selected
Viewed files
Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Unified
Split
Hide whitespace
Diff view
Unified
Split
Hide whitespace
83 changes: 73 additions & 10 deletions lib/matplotlib/figure.py
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters. Learn more about bidirectional Unicode characters
Original file line number Diff line number Diff line change
Expand Up @@ -2117,8 +2117,18 @@ def subplot_mosaic(self, mosaic, *, sharex=False, sharey=False,

per_subplot_kw = self._norm_per_subplot_kw(per_subplot_kw)

# Only accept strict bools to allow a possible future API expansion.
_api.check_isinstance(bool, sharex=sharex, sharey=sharey)
# Custom Validation Function to check for booleans and strings.
def _validate_share_param(param_name, value):
if type(value) is bool:
return value
if value in ["all", "row", "col"]:
return value
raise ValueError(
f"{param_name} must be True, False, 'all', 'row', or 'col'"
)

sharex = _validate_share_param("sharex", sharex)
sharey = _validate_share_param("sharey", sharey)

def _make_array(inp):
"""
Expand Down Expand Up @@ -2277,14 +2287,67 @@ def _do_layout(gs, mosaic, unique_ids, nested):
rows, cols = mosaic.shape
gs = self.add_gridspec(rows, cols, **gridspec_kw)
ret = _do_layout(gs, mosaic, *_identify_keys_and_nested(mosaic))
ax0 = next(iter(ret.values()))
for ax in ret.values():
if sharex:
ax.sharex(ax0)
ax._label_outer_xaxis(skip_non_rectangular_axes=True)
if sharey:
ax.sharey(ax0)
ax._label_outer_yaxis(skip_non_rectangular_axes=True)
# Asymmetrical/Directional Implementation that needs to be changed
# ax0 = next(iter(ret.values()))
# for ax in ret.values():
# if sharex:
# ax.sharex(ax0)
# ax._label_outer_xaxis(skip_non_rectangular_axes=True)
# if sharey:
# ax.sharey(ax0)
# ax._label_outer_yaxis(skip_non_rectangular_axes=True)

# Symmetric Sharing
def _apply_sharing(ret_dict, share_val, axis_name):
if share_val is False:
return

# Group the axes based on their spans
groups = {}
for ax in ret_dict.values():
span = ax.get_subplotspec()
grid = span.get_gridspec()

if share_val is True or share_val == "all":
groups.setdefault("all", []).append(ax)
elif share_val == "row":
groups.setdefault(
(grid, span.rowspan.start, span.rowspan.stop),
[]
).append(ax)
elif share_val == "col":
groups.setdefault(
(grid, span.colspan.start, span.colspan.stop),
[]
).append(ax)

# Bind the groups together
for group in groups.values():
if len(group) > 1:
parent = group[0]
for child_ax in group[1:]:
if axis_name == 'x':
child_ax.sharex(parent)
else:
child_ax.sharey(parent)

# Handle tick label visibility per isolated group
if axis_name == 'x':
# Find the bottom-most edge within a specific group
bottom_edge = max(ax.get_subplotspec().rowspan.stop for ax in group)
for ax in group:
if ax.get_subplotspec().rowspan.stop < bottom_edge:
ax.tick_params(labelbottom=False)
else:
# Find the left-most edge within a specific group
left_edge = min(ax.get_subplotspec().colspan.start for ax in group)
for ax in group:
if ax.get_subplotspec().colspan.start > left_edge:
ax.tick_params(labelleft=False)

_apply_sharing(ret, sharex, 'x')
_apply_sharing(ret, sharey, 'y')

if extra := set(per_subplot_kw) - set(ret):
raise ValueError(
f"The keys {extra} are in *per_subplot_kw* "
Expand Down
12 changes: 6 additions & 6 deletions lib/matplotlib/figure.pyi
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters. Learn more about bidirectional Unicode characters
Original file line number Diff line number Diff line change
Expand Up @@ -255,8 +255,8 @@ class FigureBase(Artist):
self,
mosaic: str,
*,
sharex: bool = ...,
sharey: bool = ...,
sharex: bool | Literal["all", "row", "col"] = ...,
sharey: bool | Literal["all", "row", "col"] = ...,
width_ratios: ArrayLike | None = ...,
height_ratios: ArrayLike | None = ...,
empty_sentinel: str = ...,
Expand All @@ -269,8 +269,8 @@ class FigureBase(Artist):
self,
mosaic: list[HashableList[T]],
*,
sharex: bool = ...,
sharey: bool = ...,
sharex: bool | Literal["all", "row", "col"] = ...,
sharey: bool | Literal["all", "row", "col"] = ...,
width_ratios: ArrayLike | None = ...,
height_ratios: ArrayLike | None = ...,
empty_sentinel: T = ...,
Expand All @@ -283,8 +283,8 @@ class FigureBase(Artist):
self,
mosaic: list[HashableList[Hashable]],
*,
sharex: bool = ...,
sharey: bool = ...,
sharex: bool | Literal["all", "row", "col"] = ...,
sharey: bool | Literal["all", "row", "col"] = ...,
width_ratios: ArrayLike | None = ...,
height_ratios: ArrayLike | None = ...,
empty_sentinel: Any = ...,
Expand Down
20 changes: 10 additions & 10 deletions lib/matplotlib/pyplot.py
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters. Learn more about bidirectional Unicode characters
Original file line number Diff line number Diff line change
Expand Up @@ -54,7 +54,7 @@
import sys
import threading
import time
from typing import IO, TYPE_CHECKING, cast, overload
from typing import IO, TYPE_CHECKING, cast, overload, Literal

from cycler import cycler # noqa: F401
import matplotlib
Expand Down Expand Up @@ -96,7 +96,7 @@
from collections.abc import Callable, Hashable, Iterable, Sequence
import pathlib
import os
from typing import Any, BinaryIO, Literal
from typing import Any, BinaryIO

import PIL.Image
from numpy.typing import ArrayLike
Expand Down Expand Up @@ -1896,8 +1896,8 @@ def subplots(
def subplot_mosaic(
mosaic: str,
*,
sharex: bool = ...,
sharey: bool = ...,
sharex: bool | Literal["all", "row", "col"] = ...,
sharey: bool | Literal["all", "row", "col"] = ...,
width_ratios: ArrayLike | None = ...,
height_ratios: ArrayLike | None = ...,
empty_sentinel: str = ...,
Expand All @@ -1912,8 +1912,8 @@ def subplot_mosaic(
def subplot_mosaic[T](
mosaic: list[HashableList[T]],
*,
sharex: bool = ...,
sharey: bool = ...,
sharex: bool | Literal["all", "row", "col"] = ...,
sharey: bool | Literal["all", "row", "col"] = ...,
width_ratios: ArrayLike | None = ...,
height_ratios: ArrayLike | None = ...,
empty_sentinel: T = ...,
Expand All @@ -1928,8 +1928,8 @@ def subplot_mosaic[T](
def subplot_mosaic(
mosaic: list[HashableList[Hashable]],
*,
sharex: bool = ...,
sharey: bool = ...,
sharex: bool | Literal["all", "row", "col"] = ...,
sharey: bool | Literal["all", "row", "col"] = ...,
width_ratios: ArrayLike | None = ...,
height_ratios: ArrayLike | None = ...,
empty_sentinel: Any = ...,
Expand All @@ -1943,8 +1943,8 @@ def subplot_mosaic(
def subplot_mosaic[T](
mosaic: str | list[HashableList[T]] | list[HashableList[Hashable]],
*,
sharex: bool = False,
sharey: bool = False,
sharex: bool | Literal["all", "row", "col"] = False,
sharey: bool | Literal["all", "row", "col"] = False,
width_ratios: ArrayLike | None = None,
height_ratios: ArrayLike | None = None,
empty_sentinel: Any = '.',
Expand Down
60 changes: 58 additions & 2 deletions lib/matplotlib/tests/test_figure.py
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters. Learn more about bidirectional Unicode characters
Original file line number Diff line number Diff line change
Expand Up @@ -1320,7 +1320,8 @@ def test_nested_user_order(self):
assert list(ax_dict) == list("ABCDEFGHI")
assert list(fig.axes) == list(ax_dict.values())

def test_share_all(self):
@pytest.mark.parametrize("share_val", [True, "all"])
def test_share_all(self, share_val):
layout = [
["A", [["B", "C"],
["D", "E"]]],
Expand All @@ -1329,11 +1330,66 @@ def test_share_all(self):
["."]]]]]
]
fig = plt.figure()
ax_dict = fig.subplot_mosaic(layout, sharex=True, sharey=True)
ax_dict = fig.subplot_mosaic(layout, sharex=share_val, sharey=share_val)
ax_dict["A"].set(xscale="log", yscale="logit")
assert all(ax.get_xscale() == "log" and ax.get_yscale() == "logit"
for ax in ax_dict.values())

def test_share_row_col(self):
layout = [
["A", "B", "C", "D"],
["E", "F", "C", "D"]
]
fig = plt.figure()
axd = fig.subplot_mosaic(layout, sharey="row", sharex="col")

# Testing row sharing for y-axis
axd["A"].set_ylim(0, 50)
axd["C"].set_ylim(-10, 10)

assert axd["B"].get_ylim() == (0, 50)
assert axd["D"].get_ylim() == (-10, 10)
assert axd["E"].get_ylim() != (0, 50)

axd["A"].set_xlim(-5, 5)
axd["C"].set_xlim(100, 200)

assert axd["E"].get_xlim() == (-5, 5)
assert axd["D"].get_xlim() != (100, 200)

def test_share_invalid(self):
layout = [
["A", "B"],
["C", "D"]
]
fig = plt.figure()

msg = "must be True, False, 'all', 'row', or 'col'"
with pytest.raises(ValueError, match=msg):
fig.subplot_mosaic(layout, sharex="invalid_string")

with pytest.raises(ValueError, match=msg):
fig.subplot_mosaic(layout, sharey={"A": "B"})

def test_share_row_col_nested(self):
layout = [
["A", [["B", "C"],
["D", "E"]]]
]
fig = plt.figure()
axd = fig.subplot_mosaic(layout, sharey="row", sharex="col")

axd["B"].set_ylim(0, 50)
assert axd["C"].get_ylim() == (0, 50)
assert axd["D"].get_ylim() != (0, 50)

axd["B"].set_xlim(0, 50)
assert axd["D"].get_xlim() == (0, 50)
assert axd["C"].get_xlim() != (0, 50)

axd["A"].set_ylim(-10, 10)
assert axd["B"].get_ylim() != (-10, 10)


def test_reused_gridspec():
"""Test that these all use the same gridspec"""
Expand Down
Loading

Back | FazBrowse Home | New Git URL