diff --git a/src/lib/data/ensure_derived.py b/src/lib/data/ensure_derived.py index d55e517..66a7dba 100644 --- a/src/lib/data/ensure_derived.py +++ b/src/lib/data/ensure_derived.py @@ -18,7 +18,7 @@ def ensure_derived[D: DataWithAttrs](data: D, key: SubdataKey) -> D: def get_derivable_keys(data: DataWithAttrs) -> list[SubdataKey]: _, prefix = split_prepath(data.metadata.prepath) if isinstance(data, Field): - return list(DERIVED_FIELD_VARIABLES[prefix].keys()) + return list(DERIVED_FIELD_VARIABLES.get(prefix, {}).keys()) elif isinstance(data, List): return list(get_derived_particle_variables(prefix).keys()) else: diff --git a/src/lib/plotting/setup_fig.py b/src/lib/plotting/setup_fig.py index 4c09120..2a5f1bf 100644 --- a/src/lib/plotting/setup_fig.py +++ b/src/lib/plotting/setup_fig.py @@ -378,9 +378,10 @@ def setup_data(self): def setup_fig(plot_infos: list[PlotInfo]) -> tuple[Figure, list[Renderer2]]: figure = plt.figure(layout="constrained") - renderers = [] + renderers: list[Renderer2] = [] - for ax, infos in _setup_axes(figure, plot_infos).values(): + loc_to_ax = _setup_axes(figure, plot_infos) + for ax, infos in loc_to_ax.values(): manager: AxesManager if len(infos) == 1: info = infos[0] @@ -407,4 +408,13 @@ def setup_fig(plot_infos: list[PlotInfo]) -> tuple[Figure, list[Renderer2]]: manager.setup() renderers += manager.renderers + # lift labels to title + if len(loc_to_ax) > 1: + suptitle_labeler = TreeLabeler(figure.suptitle("").set_text) + for renderer in renderers: + if isinstance(renderer, TreeLabeler): + suptitle_labeler.add_child(renderer) + renderers.append(suptitle_labeler) + suptitle_labeler.update() + return figure, renderers diff --git a/tests/baseline/test_suptitle.png b/tests/baseline/test_suptitle.png new file mode 100644 index 0000000..65b4b39 Binary files /dev/null and b/tests/baseline/test_suptitle.png differ diff --git a/tests/test_plots.py b/tests/test_plots.py index 0d92386..8b84b05 100644 --- a/tests/test_plots.py +++ b/tests/test_plots.py @@ -100,6 +100,12 @@ def test_image_and_cuts(): return make_plot("pfd ey_ec --scale symlog#.001 -v t y --copy ey_ec -i y=0 -v t --copy ey_ec -i y=-1 -v t".split()) +@pytest.mark.mpl_image_compare(**MPL_KWARGS) +def test_suptitle(): + """Particle densities for electrons and ions in the same figure.""" + return make_plot("prt.i -i t=1: --bin y z -v y z --with prt.e -i t=1: --bin y z -v y z loc=1,2".split()) + + # --- Cross-dataset plots ---