diff --git a/.github/workflows/ci_pipeline.yml b/.github/workflows/ci_pipeline.yml
index eced47c..d3a6406 100644
--- a/.github/workflows/ci_pipeline.yml
+++ b/.github/workflows/ci_pipeline.yml
@@ -25,9 +25,19 @@ jobs:
# with:
# python-version: ${{ matrix.python-version }}
+ - name: Install LaTeX
+ run: |
+ # pdflatex with pgfplots: the tikzfigure backend, its tests and tutorials
+ sudo apt-get update
+ sudo apt-get install -y --no-install-recommends \
+ texlive-latex-base texlive-latex-extra texlive-pictures \
+ texlive-fonts-recommended lmodern
+
- name: Install python dependencies
run: |
python -m pip install --upgrade pip
+ # tikzfigure 0.4.0 is not on PyPI yet: install it from GitHub (drop once released)
+ pip install "tikzfigure[vis] @ git+https://github.com/max-models/tikzfigure@c290f43ac1deb4a4ed3b5b40919f10d3a6b53e8b"
pip install ".[dev]"
- name: Run tests
@@ -37,7 +47,5 @@ jobs:
- name: Test tutorials
run: |
- # tutorial_07_tikz.ipynb requires pdflatex — skip it in CI
- jupyter nbconvert --to notebook --execute \
- $(ls tutorials/*.ipynb | grep -v tutorial_07_tikz) \
+ jupyter nbconvert --to notebook --execute tutorials/*.ipynb \
--output-dir=/tmp --ExecutePreprocessor.timeout=300
diff --git a/.github/workflows/docs.yml b/.github/workflows/docs.yml
index c59b8ea..8a9f6c1 100644
--- a/.github/workflows/docs.yml
+++ b/.github/workflows/docs.yml
@@ -25,14 +25,20 @@ jobs:
- name: Checkout
uses: actions/checkout@v4
- - name: Install pandoc
+ - name: Install pandoc and LaTeX
run: |
sudo apt-get update
sudo apt-get install -y pandoc
+ # pdflatex with pgfplots: the tutorials compile tikzfigure figures
+ sudo apt-get install -y --no-install-recommends \
+ texlive-latex-base texlive-latex-extra texlive-pictures \
+ texlive-fonts-recommended lmodern
- name: Install Python dependencies
run: |
python -m pip install --upgrade pip
+ # tikzfigure 0.4.0 is not on PyPI yet: install it from GitHub (drop once released)
+ pip install "tikzfigure[vis] @ git+https://github.com/max-models/tikzfigure@c290f43ac1deb4a4ed3b5b40919f10d3a6b53e8b"
pip install ".[docs]"
- name: Build Sphinx docs
diff --git a/.github/workflows/matplotlib-import.yml b/.github/workflows/matplotlib-import.yml
new file mode 100644
index 0000000..cde5a4c
--- /dev/null
+++ b/.github/workflows/matplotlib-import.yml
@@ -0,0 +1,25 @@
+name: Matplotlib import compatibility
+
+on:
+ push:
+ branches: [main, devel]
+ pull_request:
+
+jobs:
+ import-fixtures:
+ runs-on: ubuntu-latest
+ strategy:
+ fail-fast: false
+ matrix:
+ matplotlib: ['3.8.*', '3.9.*', '3.10.*', '3.11.*']
+ env:
+ MPLBACKEND: Agg
+ steps:
+ - uses: actions/checkout@v4
+ - uses: actions/setup-python@v5
+ with:
+ python-version: '3.11'
+ # tikzfigure 0.4.0 is not on PyPI yet: install it from GitHub (drop once released)
+ - run: python -m pip install "tikzfigure[vis] @ git+https://github.com/max-models/tikzfigure@c290f43ac1deb4a4ed3b5b40919f10d3a6b53e8b"
+ - run: python -m pip install '.[test]' 'numpy<2' 'matplotlib==${{ matrix.matplotlib }}'
+ - run: python -m pytest src/maxplotlib/tests/test_matplotlib_import.py src/maxplotlib/tests/test_matplotlib_import_extended.py
diff --git a/README.md b/README.md
index cbf1bcc..11f4beb 100644
--- a/README.md
+++ b/README.md
@@ -1,4 +1,4 @@
-# Maxlotlib
+# Maxplotlib
# Maxplotlib
@@ -219,23 +219,29 @@ canvas.show(backend="tikzfigure")

-### Horizontal Subplots with TikZ Backend
+### Subplots and Meshes with the TikZ Backend
-The tikzfigure backend supports creating side-by-side subplots (1×n
-layouts):
+The tikzfigure backend draws the canvas with Matplotlib and converts the
+drawn figure into pgfplots axes, so every layout converts (rows, columns,
+grids, twin axes), with LaTeX text, legends and colorbars. Lines,
+markers, bars and text are pgfplots code; meshes and images are included
+as images:
``` python
x = np.linspace(0, 2 * np.pi, 200)
-canvas, (ax1, ax2) = Canvas.subplots(ncols=2, width="10cm", ratio=0.3)
+canvas, (ax1, ax2) = Canvas.subplots(ncols=2, width="12cm", ratio=0.45)
-ax1.plot(x, np.sin(x), color="royalblue")
-ax1.set_title("sin(x)")
+ax1.plot(x, np.sin(x), color="royalblue", label="$\\sin x$")
+ax1.plot(x, np.cos(x), color="tomato", label="$\\cos x$")
+ax1.set_title("Lines")
+ax1.set_legend(True)
-ax2.plot(x, np.cos(x), color="tomato")
-ax2.set_title("cos(x)")
+xx, yy = np.meshgrid(x, x)
+ax2.pcolormesh(xx, yy, np.sin(xx) * np.cos(yy), cmap="RdBu_r")
+ax2.add_colorbar(label="$\\sin x \\cos y$")
+ax2.set_title("A mesh")
-canvas.suptitle("Trigonometric Functions")
-canvas.show(backend="tikzfigure") # Generates LaTeX subfigures
+canvas.show(backend="tikzfigure") # compiles with pdflatex
```
@@ -248,9 +254,11 @@ Figure 2
-**Note:** Only horizontal layouts (1×n) are currently supported with the
-tikzfigure backend. Vertical/grid layouts will raise
-`NotImplementedError`. See the tutorials for more examples.
+`canvas.render(backend="tikzfigure").savefig("figure.tikz")` writes the
+code for `\\input` in a LaTeX document, with the images next to it. Any
+Matplotlib figure converts the same way with
+`maxplotlib.backends.tikzfigure.figure_to_tikz(fig)`. See the tutorials
+for more examples.
### Terminal Backend with plotext
@@ -349,3 +357,49 @@ canvas.show()
(,
array([[]], dtype=object))
+
+### xarray data
+
+Plot labelled [xarray](https://docs.xarray.dev) data directly
+(`pip install maxplotlibx[xarray]`). Axes come from the coordinates,
+labels from the `long_name` and `units` attributes, and titles from the
+coordinates you selected. `import maxplotlib.xarray` adds a `.maxplot`
+accessor that mirrors xarray’s own `.plot` API and returns an ordinary
+`Canvas`, so the backend is still chosen when rendering:
+
+``` python
+import xarray as xr
+
+import maxplotlib.xarray # registers da.maxplot and ds.maxplot
+
+t = np.linspace(0, 1.5, 6)
+xs = np.linspace(0, 2 * np.pi, 80)
+ys = np.linspace(-1, 1, 50)
+wave = xr.DataArray(
+ np.sin(xs - 2 * t[:, None, None]) * np.exp(-3 * ys[None, :, None] ** 2),
+ dims=("t", "y", "x"),
+ coords={"t": ("t", t, {"units": "s"}), "y": ys, "x": ("x", xs, {"units": "m"})},
+ name="phi",
+ attrs={"long_name": "Potential", "units": "V"},
+)
+
+wave.maxplot.pcolormesh(col="t", col_wrap=3, canvas_kwargs={"width": "16cm", "ratio": 0.6}).show()
+```
+
+
+
+ (,
+ array([[,
+ ,
+ ],
+ [,
+ ,
+ ]],
+ dtype=object))
+
+The same works through Canvas methods,
+e.g. `canvas.plot(da, hue="species")`,
+`ax.pcolormesh(da, xcoord="R", ycoord="Z")` for curvilinear grids, or
+`Canvas.facet(da, col="t")`. `ds.maxplot.scatter(x=..., y=..., hue=...)`
+plots one Dataset variable against another. See the [xarray
+tutorial](tutorials/tutorial_17_xarray.ipynb) for more.
diff --git a/README.qmd b/README.qmd
index a27e50c..13111d0 100644
--- a/README.qmd
+++ b/README.qmd
@@ -1,5 +1,5 @@
---
-title: Maxlotlib
+title: Maxplotlib
format: gfm
fig-dpi: 150
---
@@ -194,9 +194,12 @@ Or plot with the TikZ backend:
canvas.show(backend="tikzfigure")
```
-### Horizontal Subplots with TikZ Backend
+### Subplots and Meshes with the TikZ Backend
-The tikzfigure backend supports creating side-by-side subplots (1×n layouts):
+The tikzfigure backend draws the canvas with Matplotlib and converts the drawn figure into
+pgfplots axes, so every layout converts (rows, columns, grids, twin axes), with LaTeX text,
+legends and colorbars. Lines, markers, bars and text are pgfplots code; meshes and images are
+included as images:
```{python}
#| label: fig-showcase-subplots
@@ -204,19 +207,24 @@ The tikzfigure backend supports creating side-by-side subplots (1×n layouts):
#| fig-height: 6
x = np.linspace(0, 2 * np.pi, 200)
-canvas, (ax1, ax2) = Canvas.subplots(ncols=2, width="10cm", ratio=0.3)
+canvas, (ax1, ax2) = Canvas.subplots(ncols=2, width="12cm", ratio=0.45)
-ax1.plot(x, np.sin(x), color="royalblue")
-ax1.set_title("sin(x)")
+ax1.plot(x, np.sin(x), color="royalblue", label="$\\sin x$")
+ax1.plot(x, np.cos(x), color="tomato", label="$\\cos x$")
+ax1.set_title("Lines")
+ax1.set_legend(True)
-ax2.plot(x, np.cos(x), color="tomato")
-ax2.set_title("cos(x)")
+xx, yy = np.meshgrid(x, x)
+ax2.pcolormesh(xx, yy, np.sin(xx) * np.cos(yy), cmap="RdBu_r")
+ax2.add_colorbar(label="$\\sin x \\cos y$")
+ax2.set_title("A mesh")
-canvas.suptitle("Trigonometric Functions")
-canvas.show(backend="tikzfigure") # Generates LaTeX subfigures
+canvas.show(backend="tikzfigure") # compiles with pdflatex
```
-**Note:** Only horizontal layouts (1×n) are currently supported with the tikzfigure backend. Vertical/grid layouts will raise `NotImplementedError`. See the tutorials for more examples.
+`canvas.render(backend="tikzfigure").savefig("figure.tikz")` writes the code for `\\input` in a
+LaTeX document, with the images next to it. Any Matplotlib figure converts the same way with
+`maxplotlib.backends.tikzfigure.figure_to_tikz(fig)`. See the tutorials for more examples.
### Terminal Backend with plotext
@@ -276,3 +284,37 @@ Show all layers:
```{python}
canvas.show()
```
+
+### xarray data
+
+Plot labelled [xarray](https://docs.xarray.dev) data directly
+(`pip install maxplotlibx[xarray]`). Axes come from the coordinates, labels
+from the `long_name` and `units` attributes, and titles from the coordinates
+you selected. `import maxplotlib.xarray` adds a `.maxplot` accessor that mirrors
+xarray's own `.plot` API and returns an ordinary `Canvas`, so the backend is
+still chosen when rendering:
+
+```{python}
+import xarray as xr
+
+import maxplotlib.xarray # registers da.maxplot and ds.maxplot
+
+t = np.linspace(0, 1.5, 6)
+xs = np.linspace(0, 2 * np.pi, 80)
+ys = np.linspace(-1, 1, 50)
+wave = xr.DataArray(
+ np.sin(xs - 2 * t[:, None, None]) * np.exp(-3 * ys[None, :, None] ** 2),
+ dims=("t", "y", "x"),
+ coords={"t": ("t", t, {"units": "s"}), "y": ys, "x": ("x", xs, {"units": "m"})},
+ name="phi",
+ attrs={"long_name": "Potential", "units": "V"},
+)
+
+wave.maxplot.pcolormesh(col="t", col_wrap=3, canvas_kwargs={"width": "16cm", "ratio": 0.6}).show()
+```
+
+The same works through Canvas methods, e.g. `canvas.plot(da, hue="species")`,
+`ax.pcolormesh(da, xcoord="R", ycoord="Z")` for curvilinear grids, or
+`Canvas.facet(da, col="t")`. `ds.maxplot.scatter(x=..., y=..., hue=...)` plots
+one Dataset variable against another. See the
+[xarray tutorial](tutorials/tutorial_17_xarray.ipynb) for more.
diff --git a/README_files/figure-commonmark/cell-20-output-1.png b/README_files/figure-commonmark/cell-20-output-1.png
new file mode 100644
index 0000000..74215c7
Binary files /dev/null and b/README_files/figure-commonmark/cell-20-output-1.png differ
diff --git a/README_files/figure-commonmark/fig-showcase-subplots-output-1.png b/README_files/figure-commonmark/fig-showcase-subplots-output-1.png
index e943f76..e37a7ac 100644
Binary files a/README_files/figure-commonmark/fig-showcase-subplots-output-1.png and b/README_files/figure-commonmark/fig-showcase-subplots-output-1.png differ
diff --git a/docs/matplotlib-import-support.md b/docs/matplotlib-import-support.md
new file mode 100644
index 0000000..5bef77d
--- /dev/null
+++ b/docs/matplotlib-import-support.md
@@ -0,0 +1,209 @@
+# Matplotlib import support and roadmap
+
+`Canvas.from_matplotlib(source, strict=False, **canvas_kwargs)` accepts an
+in-memory Matplotlib figure, axes, or rectangular axes array. This checklist
+covers the intended import surface. It is a roadmap, not a promise that every
+item already works on every rendering backend.
+
+Checked items are implemented within their stated scope. Unless explicitly
+stated otherwise, fidelity checks below refer to the **Matplotlib backend**.
+Complex built-in artists are retained as detached, editable native geometry;
+these entries deliberately raise on backends that cannot represent them.
+Portable entries continue to support the other renderers, with the limitations
+listed below. Unchecked items are
+pending or partial. Importing drawn geometry does not recover the original
+samples, plotting call, callback, or statistical model. Matplotlib, Plotly,
+TikZ, and plotext have different rendering capabilities; import success does
+not guarantee that every backend can render the result identically.
+
+## Inputs and ownership
+
+- [x] Whole `Figure`, individual `Axes`, flat lists/arrays, and 2D lists/arrays.
+- [x] Explicit 2D array order; infer ordinary grids from flat input.
+- [x] Reject empty, ragged, non-axes, duplicate, and mixed-figure inputs.
+- [x] Snapshot plot arrays and styles without reparenting source artists.
+- [x] Canvas keyword overrides for figure size, DPI, and other canvas options.
+- [x] Warnings for recognized unsupported content; strict mode raises.
+- [x] Structured import report with artist identities, severity, and fallbacks.
+- [ ] Complete detection of unsupported style properties in strict mode.
+- [ ] Version compatibility matrix and fixtures across supported Matplotlib versions.
+- [x] Optional serialized figure input with an explicit trust boundary.
+
+PNG/JPEG input and reconstruction of original data from PDF/SVG are separate
+features, not part of this object importer. Arbitrary custom Python artists
+need an adapter API or an explicit raster fallback; universal semantic
+conversion is not possible.
+
+## Figure and layout
+
+- [x] Ordinary rectangular subplot grids and individual axes.
+- [x] Figure dimensions and export DPI defaults.
+- [x] Figure title and shared axis label text.
+- [x] One `twinx` and one `twiny` per primary subplot; import only selected axes.
+- [ ] Full twin-axis spine positioning, multiple twins, and cross-backend parity.
+ Matplotlib supports multiple twins and spine positions; backend parity is pending.
+- [x] Shared-axis relationships and linked limits after import.
+- [x] Preserve omitted/empty grid cells without drawing extra empty axes.
+- [x] Spanning cells, subplot mosaics, nested GridSpec, and subfigures
+ (subfigure rectangles/decorations are flattened into the destination figure).
+- [x] Grid width/height ratios, margins, spacing, constrained/tight layout
+ for ordinary/nested GridSpec; subfigures retain their snapshot rectangles.
+- [x] Arbitrary axes rectangles, overlapping axes, inset axes, inset-zoom
+ connectors (`indicate_inset_zoom`/`indicate_inset`, Matplotlib 3.10+). The
+ connector spans two Axes and recomputes its geometry from live limits on
+ every draw, so it is rebuilt against the reconstructed parent/inset pair
+ rather than snapshotted. Matplotlib < 3.10's tuple-returning API still falls
+ back to native geometry.
+- [x] Secondary axes with forward/inverse coordinate functions.
+- [x] Figure backgrounds, frame styling, and figure-level text/patches/images.
+- [x] Figure title/shared-label typography and placement.
+- [x] Figure-level legends and shared colorbars.
+
+Irregular grids use one row of addressable Canvas slots and retain the source
+axes rectangles when rendered with Matplotlib. A 2D input array defines its own
+layout; slots occupied only by a twin remain empty rather than drawing extra axes.
+Ordinary/nested GridSpec layout engines reflow on resize. Subfigure layout-engine
+reflow and full cross-backend layout parity remain pending.
+
+## Lines, markers, and collections
+
+- [x] Numeric 2D lines, markers, NaN gaps, line/marker colors and widths.
+- [x] Step draw styles in imported line entries (Matplotlib rendering).
+- [x] Horizontal/vertical reference lines with fractional extents (Matplotlib).
+- [x] Plain line collections, including `hlines`/`vlines`, as individual lines.
+- [x] Standard scatter marker shapes, sizes, scalar colors, and per-point colors.
+- [x] Scatter colormaps and normalization copied for Matplotlib rendering.
+- [x] Preserve date/time, categorical, quantity/unit converters and formatters.
+- [x] Custom dash sequences, cap/join styles, markevery, and gap colors.
+- [x] Infinite `axline` semantics, event plots, stem containers.
+- [x] Scalar-mapped line collections, offset collections, per-point transforms.
+- [ ] Custom scatter paths and hollow markers on every backend.
+- [ ] Match scatter area/size semantics and normalization on every backend.
+- [x] Preserve ordering between artist types at equal z-order.
+
+## Bars, errors, fills, and statistical plots
+
+- [x] Vertical/horizontal bars, positions, dimensions, baselines, and styling.
+- [x] Stacked/grouped bars as their resolved rectangles.
+- [x] Histogram bars as geometry (original samples are unavailable).
+- [x] Error-bar containers with data lines: symmetric/asymmetric x/y errors,
+ caps, colors, widths, and labels; avoid duplicated component artists.
+- [x] Error-only plots and bar errors as line/cap geometry, without inventing
+ asymmetric centers that the source no longer retains.
+- [x] Simple data-coordinate polygon collections, including `fill_between`,
+ `fill_betweenx`, and stackplot regions, as filled polygon geometry.
+- [x] Data-coordinate polygon patches and stairs values/edges/baseline.
+- [x] Error limits/arrows, subsampled errors, and independently styled components.
+- [ ] Restore semantic fill boundaries/where masks when recoverable.
+ Drawn paths and retained container metadata are preserved; Matplotlib usually
+ does not retain the original where-mask or samples.
+- [x] Spans in blended coordinates; general rectangles, circles, ellipses, wedges.
+- [x] Compound polygons, holes, curved paths, arbitrary PathPatch geometry.
+- [x] Box plots, violin plots, pie/donut charts, hist2d, hexbin.
+- [x] Preserve statistical groupings when containers/metadata retain them:
+ `subplot.import_groups` and entry `source_container_ids` retain bar, errorbar,
+ and stem memberships, labels, orientation, and available data values.
+
+## Images, fields, and colorbars
+
+- [x] `imshow` arrays, extent, origin, colormap, normalization, interpolation
+ copied into entries (Matplotlib rendering).
+- [ ] Match image extent/origin, RGB(A), masks, alpha, and interpolation on all backends.
+- [ ] Nonlinear normalization, clim, under/over/bad colors on all backends.
+- [x] `pcolor`, `pcolormesh`, nonuniform grids, QuadMesh, triangular meshes.
+- [x] Contours, filled contours, levels, labels, and contour topology.
+- [x] Quiver, barbs, streamplots, vector-field keys.
+- [x] Axes colorbars tied to the correct image/scatter/mesh mappable.
+- [x] Colorbar orientation, label, ticks, limits, extend, and normalization.
+- [x] Multiple/shared colorbars and explicit colorbar axes.
+
+## Text and annotations
+
+- [x] Data-coordinate text, font family/size/weight/style, color, alignment,
+ rotation (backend support varies).
+- [x] Data-coordinate annotations with optional arrows and independent copied
+ arrow properties; references to other patches are explicitly rejected.
+- [x] Center title and x/y labels with basic typography and label padding.
+- [x] Text boxes, multiline spacing, math/TeX fidelity, wrapping, clipping
+ through Matplotlib (TeX still requires the source environment’s TeX setup).
+- [ ] Axes/figure fractions, point/pixel offsets, blended and callable coordinates.
+- [ ] Annotation arrows with full backend parity and artist-relative coordinates.
+- [x] Left/right titles, custom title/label positions, offset text.
+- [x] AnnotationBbox, OffsetBox, tables, and other composite text artists.
+
+## Axes, ticks, grids, and legends
+
+- [x] Limits including reversed limits; default linear/log scales.
+- [x] Explicit major tick locations/labels with FixedLocator.
+- [x] Basic grid visibility, axes visibility, background, aspect, axisbelow.
+- [x] Basic per-axes legend visibility and artist labels.
+- [x] Non-default log bases; symlog, logit, asinh, function/custom scales.
+- [x] Automatic/fixed minor ticks, locator/formatter configuration and units.
+- [x] Tick placement, direction, length, width, color, rotation, and font styling.
+- [x] Per-axis major/minor grid visibility and line styling.
+- [x] Spine visibility, colors, widths, bounds, and positions.
+- [x] Autoscale flags, sticky edges, margins, adjustable/anchor/box aspect.
+- [x] Legend order, renamed labels, proxy handles, multiple legends, grouping.
+- [x] Legend location/anchor, columns, title, typography, frame and spacing.
+
+## Transforms, metadata, and advanced axes
+
+- [x] General affine/nonlinear/blended transforms, coordinate rebinding.
+- [x] Clip paths/boxes, path effects, rasterization, sketch settings, filters.
+- [x] Hidden artists retained as hidden editable entries.
+- [ ] Artist IDs, URLs, picking, metadata, and accessibility descriptions.
+- [x] Polar, geographic/custom projections, axisartist and parasite axes via
+ native axes snapshots; custom projection classes must remain available.
+- [x] 3D lines/scatter, surfaces, wireframes, collections, camera/projection.
+- [ ] Animations, widgets, callbacks, and interactive state (separate adapters).
+ Static artist state is copied; arbitrary callback closures and GUI event loops
+ cannot be reconstructed from a figure’s drawn geometry.
+- [x] Optional raster fallback for unsupported artists with explicit loss reporting.
+
+## Validation and next priorities
+
+- [x] Tests for all three input forms, source independence, strict/warning modes.
+- [x] Matplotlib reconstruction tests for supported geometry and axis settings.
+- [x] Plotly smoke/geometry tests for representative supported imports.
+- [x] Image comparisons with tolerances, plus representative portable export
+ tests for Matplotlib PNG/SVG, Plotly HTML, plotext text, and TikZ source.
+- [x] Large figures, empty/masked data, performance and memory benchmarks.
+
+New import controls:
+
+```python
+canvas = Canvas.from_matplotlib(fig, strict=True)
+report = canvas.import_report.to_dict() # identities, severity, fallback, backends
+canvas = Canvas.from_matplotlib("figure.pickle", trusted=True)
+canvas = Canvas.from_matplotlib(fig, fallback="raster")
+```
+
+`trusted=True` is required before any pickle is read. Pickles can execute code;
+load only files whose producer you trust and use a matching Matplotlib version.
+`fallback="native"` (default) retains built-in native geometry with rebound
+transforms. `fallback="skip"` warns (or raises in strict mode) for those artists.
+`fallback="raster"` snapshots the selected figure content and records the loss of
+editable data and vector geometry; it supports Matplotlib and Plotly. Callback
+functions and custom scales remain Python objects, not portable serialization.
+Strict mode checks import losses, not universal rendering parity.
+
+Remaining priorities are cross-backend fidelity, subfigure reflow, portable
+projection adapters, and live interaction adapters. Source metadata does not
+usually retain original fill masks or box/violin samples; these are imported as
+geometry without inventing lost statistical inputs.
+
+Reference APIs: [Matplotlib artists](https://matplotlib.org/stable/tutorials/artists.html),
+[containers](https://matplotlib.org/stable/api/container_api.html), and
+[annotations](https://matplotlib.org/stable/users/explain/text/annotations.html).
+
+
+Compatibility fixtures live in `test_matplotlib_import.py` and
+`test_matplotlib_import_extended.py`. The CI matrix covers Matplotlib 3.8, 3.9,
+3.10 and 3.11 on Python 3.11. Snapshots use version-sensitive Matplotlib state;
+there is no cross-version pickle compatibility promise.
+
+Run `python scripts/benchmark_matplotlib_import.py --points 100000 --panels 4`
+for reproducible timing/allocation observations. A local Matplotlib 3.11.1 run
+imported 400,000 points in 0.41 s with 8.75 MiB peak Python allocations and rendered
+in 0.79 s. These are observations, not performance guarantees; native allocations
+outside Python are excluded from the memory measurement.
diff --git a/docs/source/index.rst b/docs/source/index.rst
index e1bd3ab..82a2734 100644
--- a/docs/source/index.rst
+++ b/docs/source/index.rst
@@ -31,3 +31,4 @@ documentation for details.
tutorials/tutorial_14_axis_and_layout_controls
tutorials/tutorial_15_tikzfigure_subplots
tutorials/tutorial_16_plotext_advanced
+ tutorials/tutorial_17_xarray
diff --git a/pyproject.toml b/pyproject.toml
index e7c49cf..5c3163a 100644
--- a/pyproject.toml
+++ b/pyproject.toml
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
[project]
name = "maxplotlibx"
-version = "0.1.9"
+version = "0.2.0"
description = "A reproducible plotting module with various backends and export options."
readme = "README.md"
requires-python = ">=3.8"
@@ -19,14 +19,19 @@ dependencies = [
"pint",
"plotly",
"plotext >= 6.0, < 7",
- "tikzfigure[vis]>=0.3.0",
+ "tikzfigure[vis]>=0.4.0",
]
[project.optional-dependencies]
test = [
"pytest",
"coverage",
+ "xarray",
+]
+xarray = [
+ "xarray",
]
docs = [
+ "xarray",
"myst-parser",
"sphinx",
"sphinx-rtd-theme",
diff --git a/scripts/benchmark_matplotlib_import.py b/scripts/benchmark_matplotlib_import.py
new file mode 100644
index 0000000..1121383
--- /dev/null
+++ b/scripts/benchmark_matplotlib_import.py
@@ -0,0 +1,58 @@
+"""Measure import time and peak Python allocation for reproducible fixtures.
+
+Run from an installed checkout: python scripts/benchmark_matplotlib_import.py
+Timings are observations, not brittle pass/fail thresholds. Native renderer
+allocations outside Python are not included in tracemalloc's peak measurement.
+"""
+
+import argparse
+import json
+import time
+import tracemalloc
+
+import matplotlib
+
+matplotlib.use("Agg")
+import matplotlib.pyplot as plt
+import numpy as np
+
+from maxplotlib import Canvas
+
+
+def benchmark(points, panels):
+ fig, axes = plt.subplots(panels, 1, figsize=(8, max(3, panels * 2)))
+ x = np.linspace(0, 10, points)
+ for index, ax in enumerate(np.atleast_1d(axes)):
+ ax.plot(x, np.sin(x + index))
+ ax.scatter([], [])
+ tracemalloc.start()
+ started = time.perf_counter()
+ canvas = Canvas.from_matplotlib(fig, strict=True)
+ imported = time.perf_counter()
+ _, peak = tracemalloc.get_traced_memory()
+ tracemalloc.stop()
+ result, _ = canvas.render()
+ result.canvas.draw()
+ finished = time.perf_counter()
+ stats = dict(
+ matplotlib=matplotlib.__version__,
+ points_per_panel=points,
+ panels=panels,
+ import_seconds=imported - started,
+ render_seconds=finished - imported,
+ import_peak_python_mib=peak / 1024**2,
+ entries=sum(len(plot.line_data) for _, _, plot in canvas.iter_subplots()),
+ )
+ plt.close(fig)
+ plt.close(result)
+ return stats
+
+
+if __name__ == "__main__":
+ parser = argparse.ArgumentParser(description=__doc__)
+ parser.add_argument("--points", type=int, default=100_000)
+ parser.add_argument("--panels", type=int, default=4)
+ arguments = parser.parse_args()
+ if arguments.points < 1 or arguments.panels < 1:
+ parser.error("points and panels must be positive")
+ print(json.dumps(benchmark(arguments.points, arguments.panels), indent=2))
diff --git a/src/maxplotlib/backends/matplotlib/import_state.py b/src/maxplotlib/backends/matplotlib/import_state.py
new file mode 100644
index 0000000..286dedc
--- /dev/null
+++ b/src/maxplotlib/backends/matplotlib/import_state.py
@@ -0,0 +1,342 @@
+"""Detached Matplotlib state used by the object importer.
+
+Transforms must be rebound, not frozen in source display coordinates. The
+snapshot memo cuts ownership links before copying and replaces them with tokens;
+each render binds those tokens to its new figure/axes. No source artists are
+reparented, and repeated renders do not share mutable Matplotlib state.
+"""
+
+import copy
+from dataclasses import asdict, dataclass, field
+
+from matplotlib.artist import Artist
+from matplotlib.cbook import CallbackRegistry
+from matplotlib.collections import Collection
+from matplotlib.image import AxesImage
+from matplotlib.lines import Line2D
+from matplotlib.patches import Patch
+from matplotlib.table import Table
+from matplotlib.text import Text
+from matplotlib.transforms import Bbox, BboxTransformTo
+
+
+@dataclass
+class ImportDiagnostic:
+ artist_id: int
+ artist_type: str
+ label: str
+ severity: str
+ message: str
+ fallback: str | None = None
+ backends: tuple = ("matplotlib",)
+
+
+@dataclass
+class ImportReport:
+ """Import decisions, including losses and backend-specific representations."""
+
+ diagnostics: list = field(default_factory=list)
+
+ def add(
+ self,
+ artist,
+ message,
+ *,
+ severity="info",
+ fallback=None,
+ backends=("matplotlib",),
+ ):
+ item = ImportDiagnostic(
+ id(artist),
+ type(artist).__name__,
+ str(getattr(artist, "get_label", lambda: "")()),
+ severity,
+ message,
+ fallback,
+ backends,
+ )
+ self.diagnostics.append(item)
+ return item
+
+ def to_dict(self):
+ return asdict(self)
+
+
+def root_figure(figure):
+ while getattr(figure, "figure", figure) is not figure:
+ figure = figure.figure
+ return figure
+
+
+def figure_bounds(ax, figure):
+ return tuple(
+ ax.get_position(original=True)
+ .transformed(ax.figure.transSubfigure)
+ .transformed(figure.transFigure.inverted())
+ .bounds
+ )
+
+
+def _bindings(ax, fig):
+ bindings = {
+ "figure": fig,
+ "transFigure": fig.transFigure,
+ "transSubfigure": fig.transSubfigure,
+ "dpi_scale_trans": fig.dpi_scale_trans,
+ "fig_bbox": fig.bbox,
+ }
+ if ax is not None:
+ bindings.update(
+ axes=ax,
+ transData=ax.transData,
+ transAxes=ax.transAxes,
+ transScale=ax.transScale,
+ transLimits=ax.transLimits,
+ ax_bbox=ax.bbox,
+ xaxis=ax.xaxis,
+ yaxis=ax.yaxis,
+ xaxis_transform=ax.get_xaxis_transform(),
+ yaxis_transform=ax.get_yaxis_transform(),
+ )
+ return bindings
+
+
+class _BindingToken:
+ """Weak-referenceable marker for copied Transform parent links."""
+
+
+class ReboundSnapshot:
+ """Copy a payload without retaining its owning axes, figure or callbacks."""
+
+ def __init__(self, payload, ax=None, fig=None):
+ owner = fig if fig is not None else ax.figure
+ fig = root_figure(owner)
+ self.owner_bounds = None
+ self.tokens = {}
+ memo = {}
+ bindings = _bindings(ax, fig)
+ if owner is not fig:
+ self.owner_bounds = tuple(
+ owner.bbox.transformed(fig.transFigure.inverted()).bounds
+ )
+ bindings.update(owner=owner, owner_transform=owner.transSubfigure)
+ for name, value in bindings.items():
+ # Some canonical transforms are aliases of one another.
+ if id(value) not in memo:
+ self.tokens[name] = memo[id(value)] = _BindingToken()
+ artists = []
+ if isinstance(payload, Artist):
+ artists = payload.findobj()
+ for artist in artists:
+ for name in ("_remove_method", "stale_callback"):
+ value = getattr(artist, name, None)
+ if value is not None:
+ memo[id(value)] = None
+ callbacks = getattr(artist, "_callbacks", None)
+ if callbacks is not None:
+ memo[id(callbacks)] = CallbackRegistry()
+ self.payload = copy.deepcopy(payload, memo)
+
+ def clone(self, ax=None, fig=None):
+ fig = fig if fig is not None else ax.figure
+ bindings = _bindings(ax, fig)
+ if self.owner_bounds is not None:
+ bindings.update(
+ owner=fig,
+ owner_transform=BboxTransformTo(Bbox.from_bounds(*self.owner_bounds))
+ + fig.transFigure,
+ )
+ return copy.deepcopy(
+ self.payload,
+ {id(token): bindings[name] for name, token in self.tokens.items()},
+ )
+
+ def draw(self, ax, **overrides):
+ artist = self.clone(ax)
+ artist.set(**overrides)
+ clipbox, clippath = artist.get_clip_box(), artist.get_clip_path()
+ # add_* establishes the new removal and stale callbacks.
+ if isinstance(artist, Collection):
+ ax.add_collection(artist, autolim=False)
+ elif isinstance(artist, Line2D):
+ ax.add_line(artist)
+ elif isinstance(artist, Patch):
+ ax.add_patch(artist)
+ elif isinstance(artist, AxesImage):
+ ax.add_image(artist)
+ elif isinstance(artist, Table):
+ ax.add_table(artist)
+ elif isinstance(artist, Text):
+ ax._add_text(artist)
+ else:
+ ax.add_artist(artist)
+ artist.set_clip_path(clippath)
+ artist.set_clip_box(clipbox)
+ return artist
+
+
+def capture_axis_state(ax):
+ """Capture decorations which have no backend-neutral equivalent yet."""
+ state = {"axes": {}, "spines": {}, "titles": {}}
+ for name in ("x", "y"):
+ axis = getattr(ax, name + "axis")
+ state["axes"][name] = dict(
+ scale=ReboundSnapshot(axis._scale, ax),
+ major_locator=ReboundSnapshot(axis.get_major_locator(), ax),
+ minor_locator=ReboundSnapshot(axis.get_minor_locator(), ax),
+ major_formatter=ReboundSnapshot(axis.get_major_formatter(), ax),
+ minor_formatter=ReboundSnapshot(axis.get_minor_formatter(), ax),
+ units=copy.deepcopy(axis.get_units()),
+ converter=copy.deepcopy(
+ axis.get_converter()
+ if hasattr(axis, "get_converter")
+ else axis.converter
+ ),
+ major_kw=_tick_params(axis, "major"),
+ minor_kw=_tick_params(axis, "minor"),
+ ticks={
+ which: [
+ _tick_style(tick)
+ for tick in getattr(axis, "get_" + which + "_ticks")()
+ ]
+ for which in ("major", "minor")
+ },
+ label_position=axis.get_label_position(),
+ offset_style=_text_properties(axis.get_offset_text()),
+ )
+ for name, spine in ax.spines.items():
+ state["spines"][name] = dict(
+ visible=spine.get_visible(),
+ edgecolor=spine.get_edgecolor(),
+ linewidth=spine.get_linewidth(),
+ linestyle=spine.get_linestyle(),
+ bounds=spine.get_bounds(),
+ )
+ if spine.spine_type in ("left", "right", "top", "bottom"):
+ state["spines"][name]["position"] = copy.deepcopy(spine.get_position())
+ for loc, title in (
+ ("left", ax._left_title),
+ ("center", ax.title),
+ ("right", ax._right_title),
+ ):
+ state["titles"][loc] = (
+ title.get_text(),
+ _text_properties(title),
+ title.get_position(),
+ )
+ state["autotitlepos"] = ax._autotitlepos
+ state["label_coords"] = {
+ name: (
+ axis.label.get_position(),
+ ReboundSnapshot(axis.label.get_transform(), ax),
+ )
+ for name, axis in (("x", ax.xaxis), ("y", ax.yaxis))
+ if not axis._autolabelpos
+ }
+ return state
+
+
+def _text_properties(text):
+ return dict(
+ fontproperties=copy.deepcopy(text.get_fontproperties()),
+ color=text.get_color(),
+ rotation=text.get_rotation(),
+ rotation_mode=text.get_rotation_mode(),
+ horizontalalignment=text.get_ha(),
+ verticalalignment=text.get_va(),
+ visible=text.get_visible(),
+ usetex=text.get_usetex(),
+ )
+
+
+def _tick_params(axis, which):
+ tick = getattr(axis, "get_" + which + "_ticks")(1)[0]
+ defaults = dict(
+ length=tick._size,
+ width=tick._width,
+ pad=tick._base_pad,
+ direction=tick._tickdir,
+ color=tick.tick1line.get_color(),
+ )
+ defaults.update(copy.deepcopy(getattr(axis, "_" + which + "_tick_kw")))
+ return defaults
+
+
+def _tick_line_style(line):
+ return {
+ name: getattr(line, "get_" + name)()
+ for name in (
+ "color",
+ "marker",
+ "markersize",
+ "markeredgewidth",
+ "markeredgecolor",
+ "visible",
+ "zorder",
+ )
+ }
+
+
+def _tick_style(tick):
+ return dict(
+ tick1line=_tick_line_style(tick.tick1line),
+ tick2line=_tick_line_style(tick.tick2line),
+ label1=_text_properties(tick.label1),
+ label2=_text_properties(tick.label2),
+ gridline=dict(
+ visible=tick.gridline.get_visible(),
+ color=tick.gridline.get_color(),
+ linewidth=tick.gridline.get_linewidth(),
+ linestyle=tick.gridline.get_linestyle(),
+ alpha=tick.gridline.get_alpha(),
+ ),
+ )
+
+
+def apply_axis_state(ax, state, *, units=True):
+ for name, props in state["spines"].items():
+ props = copy.deepcopy(props)
+ bounds = props.pop("bounds")
+ ax.spines[name].set(**props)
+ if bounds is not None:
+ ax.spines[name].set_bounds(*bounds)
+ for name, settings in state["axes"].items():
+ axis = getattr(ax, name + "axis")
+ # Scales are installed before locators/formatters: changing a scale
+ # installs defaults and would otherwise erase the imported tick setup.
+ getattr(ax, "set_" + name + "scale")(settings["scale"].clone(ax))
+ if units:
+ converter = copy.deepcopy(settings["converter"])
+ if hasattr(axis, "set_converter"):
+ if not getattr(axis, "_converter_is_explicit", False):
+ axis.set_converter(converter)
+ else:
+ axis.converter = converter
+ axis.set_units(copy.deepcopy(settings["units"]))
+ for kind in ("major", "minor"):
+ getattr(axis, "set_" + kind + "_locator")(
+ settings[kind + "_locator"].clone(ax)
+ )
+ getattr(axis, "set_" + kind + "_formatter")(
+ settings[kind + "_formatter"].clone(ax)
+ )
+ axis.set_tick_params(which=kind, **copy.deepcopy(settings[kind + "_kw"]))
+ ticks = getattr(axis, "get_" + kind + "_ticks")()
+ styles = settings["ticks"][kind]
+ for i, tick in enumerate(ticks):
+ if styles:
+ style = styles[min(i, len(styles) - 1)]
+ for part, props in style.items():
+ getattr(tick, part).set(**props)
+ axis.set_label_position(settings["label_position"])
+ axis.get_offset_text().set(**settings["offset_style"])
+ for loc, (text, props, position) in (
+ state["titles"].items() if hasattr(ax, "set_title") else ()
+ ):
+ title = ax.set_title(text, loc=loc, **props)
+ title.set_position(position)
+ ax._autotitlepos = state["autotitlepos"]
+ for name, (position, transform) in state["label_coords"].items():
+ getattr(ax, name + "axis").set_label_coords(
+ *position, transform=transform.clone(ax)
+ )
diff --git a/src/maxplotlib/backends/matplotlib/importer.py b/src/maxplotlib/backends/matplotlib/importer.py
new file mode 100644
index 0000000..15036b2
--- /dev/null
+++ b/src/maxplotlib/backends/matplotlib/importer.py
@@ -0,0 +1,1078 @@
+"""Translate Matplotlib artists into backend-independent canvas entries."""
+
+import copy
+import os
+import pickle
+import warnings
+
+import numpy as np
+from matplotlib.artist import Artist
+from matplotlib.axes import Axes
+from matplotlib.collections import LineCollection, PathCollection, PolyCollection
+from matplotlib.container import BarContainer, ErrorbarContainer, StemContainer
+from matplotlib.figure import Figure
+from matplotlib.legend import Legend
+from matplotlib.lines import AxLine
+from matplotlib.markers import MarkerStyle
+from matplotlib.patches import Polygon, StepPatch
+from matplotlib.path import Path
+from matplotlib.text import Annotation
+from matplotlib.ticker import FixedLocator
+from matplotlib.transforms import IdentityTransform
+
+try:
+ from matplotlib.inset import InsetIndicator
+except ImportError: # Matplotlib < 3.10 returns a plain (Rectangle, patches) tuple.
+ InsetIndicator = None
+
+from .import_state import (
+ ImportReport,
+ ReboundSnapshot,
+ capture_axis_state,
+ figure_bounds,
+ root_figure,
+)
+
+
+def import_matplotlib(
+ canvas_cls,
+ source,
+ *,
+ strict=False,
+ trusted=False,
+ fallback="native",
+ **canvas_kwargs,
+):
+ if fallback not in ("native", "skip", "raster"):
+ raise ValueError("fallback must be 'native', 'skip', or 'raster'")
+ if isinstance(source, (str, bytes, os.PathLike)) or hasattr(source, "read"):
+ if not trusted:
+ raise ValueError(
+ "Serialized Matplotlib input requires trusted=True; "
+ "unpickling can execute arbitrary Python code"
+ )
+ if isinstance(source, bytes):
+ source = pickle.loads(source)
+ elif hasattr(source, "read"):
+ source = pickle.load(source)
+ else:
+ with open(source, "rb") as stream:
+ source = pickle.load(stream)
+ report = ImportReport()
+
+ def unsupported(message, artist=None):
+ artist = artist if artist is not None else unsupported.artist
+ report.add(
+ artist,
+ message,
+ severity="error" if strict else "warning",
+ fallback="skipped",
+ )
+ if strict:
+ error = NotImplementedError(message)
+ error.import_report = report
+ raise error
+ warnings.warn(message, UserWarning, stacklevel=3)
+
+ unsupported.artist = source
+ unsupported.report = report
+ unsupported.fallback = fallback
+
+ whole_figure = isinstance(source, Figure)
+ explicit_shape = None
+ if whole_figure:
+ axes = list(source.axes)
+ elif isinstance(source, Axes):
+ axes = [source]
+ explicit_shape = (1, 1)
+ else:
+ try:
+ array = np.asarray(source, dtype=object)
+ except (TypeError, ValueError) as exc:
+ raise TypeError(
+ "Expected a Figure, Axes, or rectangular array of Axes"
+ ) from exc
+ if array.ndim not in (1, 2):
+ raise TypeError("Expected a Figure, Axes, or 1D/2D array of Axes")
+ axes = list(array.flat)
+ if array.ndim == 2:
+ explicit_shape = array.shape
+ if not axes:
+ raise ValueError("At least one Matplotlib Axes is required")
+ if not all(isinstance(ax, Axes) for ax in axes):
+ raise TypeError("Every array element must be a Matplotlib Axes")
+ if len({id(ax) for ax in axes}) != len(axes):
+ raise ValueError("Duplicate Axes are not supported")
+ figure = root_figure(source if whole_figure else axes[0].get_figure())
+ if any(root_figure(ax.get_figure()) is not figure for ax in axes):
+ raise ValueError("All Axes must belong to the same Figure")
+
+ if fallback == "raster":
+ return _raster_import(
+ canvas_cls, figure, axes, whole_figure, report, **canvas_kwargs
+ )
+
+ # Colorbar axes are decorations, not independent data subplots.
+ selected = []
+ twins = []
+ colorbars = []
+ # The source figure's creation order identifies the primary axes even when
+ # the caller passes the twin first in an array.
+ for index, ax in sorted(
+ enumerate(axes), key=lambda item: figure.axes.index(item[1])
+ ):
+ if getattr(ax, "_colorbar", None) is not None:
+ colorbars.append(ax._colorbar)
+ else:
+ parent = next(
+ (
+ previous
+ for _, previous in selected
+ if ax in previous._twinned_axes.get_siblings(previous)
+ ),
+ None,
+ )
+ if parent is None:
+ selected.append((index, ax))
+ else:
+ direction = "x" if ax.get_shared_x_axes().joined(ax, parent) else "y"
+ twins.append((parent, ax, direction))
+ selected.sort(key=lambda item: item[0])
+ if not selected:
+ raise ValueError("No supported Axes to import")
+
+ if explicit_shape is not None:
+ nrows, ncols = explicit_shape
+ positioned = [(i // ncols, i % ncols, ax) for i, ax in selected]
+ else:
+ specs = [ax.get_subplotspec() for _, ax in selected]
+ common_grid = specs[0].get_gridspec() if specs[0] is not None else None
+ simple_grid = common_grid is not None and all(
+ spec is not None
+ and spec.get_gridspec() is common_grid
+ and spec.rowspan.stop - spec.rowspan.start == 1
+ and spec.colspan.stop - spec.colspan.start == 1
+ for spec in specs
+ )
+ if simple_grid and len({spec.num1 for spec in specs}) == len(specs):
+ nrows, ncols = common_grid.get_geometry()
+ positioned = [
+ (spec.rowspan.start, spec.colspan.start, ax)
+ for spec, (_, ax) in zip(specs, selected)
+ ]
+ else:
+ # Irregular layouts use stable editable slots and retain their
+ # actual axes rectangles for the Matplotlib renderer.
+ nrows, ncols = 1, len(selected)
+ positioned = [(0, i, ax) for i, (_, ax) in enumerate(selected)]
+
+ if "nrows" in canvas_kwargs or "ncols" in canvas_kwargs:
+ raise TypeError("nrows and ncols are inferred from the imported axes")
+ canvas_kwargs.setdefault("figsize", tuple(figure.get_size_inches()))
+ canvas_kwargs.setdefault("dpi", figure.dpi)
+ canvas = canvas_cls(nrows=nrows, ncols=ncols, **canvas_kwargs)
+ canvas.import_report = report
+ canvas._import_layout = {}
+ canvas._import_in_layout = {}
+ canvas._import_shared_axes = []
+ canvas._import_extra_twins = []
+ canvas._import_colorbars = []
+ canvas._import_figure_artists = []
+ source_targets = {}
+ source_slots = {}
+ for row, col, ax in positioned:
+ target = canvas.add_subplot(row=row, col=col)
+ source_targets[ax] = target
+ source_slots[ax] = (row, col)
+ if explicit_shape is None:
+ canvas._import_layout[(row, col)] = figure_bounds(ax, figure)
+ canvas._import_in_layout[(row, col)] = ax.get_in_layout()
+ _import_axes(ax, target, unsupported)
+ for parent, twin, direction in twins:
+ if parent is ax:
+ factory = canvas.twinx if direction == "x" else canvas.twiny
+ registry = (
+ canvas._twinx_subplots
+ if direction == "x"
+ else canvas._twiny_subplots
+ )
+ if (row, col) in registry:
+ from maxplotlib.subfigure.line_plot import LinePlot
+
+ twin_target = LinePlot()
+ canvas._import_extra_twins.append(
+ ((row, col), direction, twin_target)
+ )
+ else:
+ twin_target = factory(row=row, col=col)
+ source_targets[twin] = twin_target
+ _import_axes(twin, twin_target, unsupported)
+ if explicit_shape is None and all(ax.figure is figure for _, ax in selected):
+ specs = {source_slots[ax]: ax.get_subplotspec() for _, ax in selected}
+ canvas._import_layout_specs = ReboundSnapshot(specs, fig=figure)
+ canvas._import_layout_engine = copy.deepcopy(figure.get_layout_engine())
+ if canvas._import_layout_engine is None:
+ canvas.subplots_adjust(
+ **{
+ name: getattr(figure.subplotpars, name)
+ for name in ("left", "right", "bottom", "top", "wspace", "hspace")
+ }
+ )
+ for name in ("x", "y"):
+ for index, (_, ax) in enumerate(selected):
+ parent = next(
+ (
+ previous
+ for _, previous in selected[:index]
+ if getattr(ax, "get_shared_" + name + "_axes")().joined(
+ ax, previous
+ )
+ ),
+ None,
+ )
+ if parent is not None:
+ canvas._import_shared_axes.append(
+ (name, source_slots[parent], source_slots[ax])
+ )
+ for colorbar in colorbars:
+ owner = getattr(colorbar.mappable, "axes", None)
+ if owner is not None and owner not in source_targets:
+ unsupported("Colorbar mappable is outside the selected axes", colorbar.ax)
+ continue
+ canvas._import_colorbars.append(
+ _capture_colorbar(colorbar, source_targets.get(owner))
+ )
+ report.add(colorbar.ax, "Colorbar is linked to its imported mappable")
+ if whole_figure:
+ canvas._import_figure_patch = ReboundSnapshot(figure.patch, fig=figure)
+ canvas._import_figure_style = dict(
+ facecolor=figure.get_facecolor(),
+ edgecolor=figure.get_edgecolor(),
+ linewidth=figure.patch.get_linewidth(),
+ frameon=figure.get_frameon(),
+ )
+ for getter, setter in (
+ ("get_suptitle", "suptitle"),
+ ("get_supxlabel", "supxlabel"),
+ ("get_supylabel", "supylabel"),
+ ):
+ value = getattr(figure, getter, lambda: "")()
+ if value:
+ text = getattr(figure, "_" + setter)
+ getattr(canvas, setter)(
+ value,
+ x=text.get_position()[0],
+ y=text.get_position()[1],
+ **_text_style(text),
+ )
+ titles = {figure._suptitle, figure._supxlabel, figure._supylabel}
+ for artist in (
+ list(figure.legends)
+ + list(figure.artists)
+ + list(figure.lines)
+ + list(figure.patches)
+ + list(figure.images)
+ + [text for text in figure.texts if text not in titles]
+ ):
+ canvas._import_figure_artists.append(ReboundSnapshot(artist, fig=figure))
+ report.add(
+ artist, "Figure decoration retained for Matplotlib", fallback="native"
+ )
+
+ def subfigure_decorations(owner):
+ for subfigure in owner.subfigs:
+ background = ReboundSnapshot(subfigure.patch, fig=subfigure)
+ background.payload.set_zorder(
+ min((ax.get_zorder() for ax in figure.axes), default=0) - 1
+ )
+ canvas._import_figure_artists.append(background)
+ for artist in (
+ list(subfigure.texts)
+ + list(subfigure.legends)
+ + list(subfigure.artists)
+ + list(subfigure.lines)
+ + list(subfigure.patches)
+ + list(subfigure.images)
+ ):
+ canvas._import_figure_artists.append(
+ ReboundSnapshot(artist, fig=subfigure)
+ )
+ subfigure_decorations(subfigure)
+
+ subfigure_decorations(figure)
+ return canvas
+
+
+def _style(artist):
+ return dict(
+ label=artist.get_label(),
+ alpha=copy.deepcopy(artist.get_alpha()),
+ zorder=artist.get_zorder(),
+ visible=artist.get_visible(),
+ gid=artist.get_gid(),
+ url=artist.get_url(),
+ rasterized=artist.get_rasterized(),
+ clip_on=artist.get_clip_on(),
+ snap=artist.get_snap(),
+ in_layout=artist.get_in_layout(),
+ )
+
+
+def _scatter_marker(path):
+ # Keep standard marker names portable to the other renderers.
+ for marker in (
+ "o",
+ "s",
+ "^",
+ "v",
+ "<",
+ ">",
+ "D",
+ "d",
+ "*",
+ "+",
+ "x",
+ "p",
+ "h",
+ "H",
+ ".",
+ "P",
+ "X",
+ ):
+ style = MarkerStyle(marker)
+ candidate = style.get_path().transformed(style.get_transform())
+ if np.array_equal(path.vertices, candidate.vertices) and np.array_equal(
+ path.codes, candidate.codes
+ ):
+ return marker
+ return copy.deepcopy(path)
+
+
+def _line_style(line):
+ return dict(
+ color=line.get_color(),
+ linestyle=(
+ copy.deepcopy(line._unscaled_dash_pattern)
+ if line.is_dashed()
+ else line.get_linestyle()
+ ),
+ linewidth=line.get_linewidth(),
+ marker=None if line.get_marker() in ("None", "", " ") else line.get_marker(),
+ markersize=line.get_markersize(),
+ markerfacecolor=line.get_markerfacecolor(),
+ markeredgecolor=line.get_markeredgecolor(),
+ markeredgewidth=line.get_markeredgewidth(),
+ drawstyle=line.get_drawstyle(),
+ dash_capstyle=line.get_dash_capstyle(),
+ dash_joinstyle=line.get_dash_joinstyle(),
+ solid_capstyle=line.get_solid_capstyle(),
+ solid_joinstyle=line.get_solid_joinstyle(),
+ markevery=copy.deepcopy(line.get_markevery()),
+ gapcolor=line.get_gapcolor(),
+ antialiased=line.get_antialiased(),
+ markerfacecoloralt=line.get_markerfacecoloralt(),
+ fillstyle=line.get_fillstyle(),
+ **_style(line),
+ )
+
+
+def _text_style(text):
+ kwargs = dict(
+ color=text.get_color(),
+ fontsize=text.get_fontsize(),
+ fontweight=text.get_fontweight(),
+ fontstyle=text.get_fontstyle(),
+ fontfamily=list(text.get_fontfamily()),
+ rotation=text.get_rotation(),
+ ha=text.get_ha(),
+ va=text.get_va(),
+ usetex=text.get_usetex(),
+ rotation_mode=text.get_rotation_mode(),
+ linespacing=text._linespacing,
+ fontstretch=text.get_fontproperties().get_stretch(),
+ fontvariant=text.get_fontproperties().get_variant(),
+ parse_math=text.get_parse_math(),
+ wrap=text.get_wrap(),
+ )
+
+ box = text.get_bbox_patch()
+ if box is not None:
+ kwargs["bbox"] = dict(
+ boxstyle=copy.deepcopy(box.get_boxstyle()),
+ facecolor=box.get_facecolor(),
+ edgecolor=box.get_edgecolor(),
+ linewidth=box.get_linewidth(),
+ linestyle=box.get_linestyle(),
+ alpha=box.get_alpha(),
+ hatch=box.get_hatch(),
+ )
+ return kwargs
+
+
+def _import_errorbar(container, ax, target, unsupported):
+ """Return whether the container was consumed as one semantic entry."""
+ line, caps, ranges = container.lines
+ # Without a data line the original center of asymmetric errors is lost.
+ # Leave these artists for the geometry import instead of inventing centers.
+ if line is None:
+ return False
+ if line.get_transform() != ax.transData or any(
+ bars.get_transform() != ax.transData for bars in ranges
+ ):
+ return False
+ components = [line, *caps, *ranges]
+ if any(
+ _needs_native_style(artist) or artist.get_visible() != line.get_visible()
+ for artist in components
+ ):
+ return False
+ if any(cap.get_marker() not in ("_", "|") for cap in caps):
+ return False
+ if caps and any(
+ any(
+ getattr(cap, "get_" + prop)() != getattr(caps[0], "get_" + prop)()
+ for prop in ("color", "markersize", "markeredgewidth", "alpha", "visible")
+ )
+ for cap in caps[1:]
+ ):
+ return False
+ x, y = (np.array(v, copy=True) for v in line.get_data(orig=False))
+ errors = {}
+ for axis, bars in zip(
+ [i for i, flag in enumerate((container.has_xerr, container.has_yerr)) if flag],
+ ranges,
+ ):
+ segments = bars.get_segments()
+ if len(segments) != len(x):
+ return False
+ bounds = np.full((len(x), 2), np.nan)
+ for i, segment in enumerate(segments):
+ if len(segment) == 2:
+ bounds[i] = segment[:, axis]
+ centers = x if axis == 0 else y
+ error = np.array([centers - bounds[:, 0], bounds[:, 1] - centers])
+ if np.any(error < -1e-12):
+ unsupported(
+ "Error-bar centers are outside their ranges; importing geometry"
+ )
+ return False
+ errors["xerr" if axis == 0 else "yerr"] = np.maximum(error, 0)
+ kwargs = _line_style(line)
+ kwargs["label"] = container.get_label()
+ if any(
+ not np.array_equal(bars.get_colors(), ranges[0].get_colors())
+ or not np.array_equal(bars.get_linewidths(), ranges[0].get_linewidths())
+ for bars in ranges[1:]
+ ):
+ return False
+ if ranges:
+ base_zorder = ranges[0].get_zorder()
+ delta = line.get_zorder() - base_zorder
+ if not np.isclose(abs(delta), 0.1) or any(
+ bar.get_zorder() != base_zorder for bar in ranges
+ ):
+ return False
+ kwargs["zorder"] = base_zorder
+ kwargs["barsabove"] = delta < 0
+ colors = ranges[0].get_colors()
+ widths = ranges[0].get_linewidths()
+ if len(colors):
+ kwargs["ecolor"] = tuple(colors[0])
+ if len(widths):
+ kwargs["elinewidth"] = float(widths[0])
+ kwargs["capsize"] = caps[0].get_markersize() / 2 if caps else 0
+ if caps:
+ kwargs["capthick"] = caps[0].get_markeredgewidth()
+ target.errorbar(x, y, **errors, **kwargs)
+ return True
+
+
+def _collection_style(collection, index):
+ kwargs = _style(collection)
+ if index:
+ kwargs["label"] = ""
+ for key, values in (
+ ("facecolor", collection.get_facecolors()),
+ ("edgecolor", collection.get_edgecolors()),
+ ("linewidth", collection.get_linewidths()),
+ (
+ "linestyle",
+ getattr(collection, "_us_linestyles", collection.get_linestyles()),
+ ),
+ ("antialiased", collection._antialiaseds),
+ ):
+ kwargs[key] = (
+ copy.deepcopy(values[index % len(values)]) if len(values) else "none"
+ )
+ offset, dashes = kwargs["linestyle"]
+ kwargs["linestyle"] = (offset, tuple(dashes) if dashes is not None else None)
+ return kwargs
+
+
+def _import_collection(collection, ax, target, unsupported):
+ """Import plain line/polygon collections; return False for scatter."""
+ if not isinstance(collection, (LineCollection, PolyCollection)):
+ return False
+ if (
+ collection.get_transform() != ax.transData
+ or np.any(collection.get_offsets())
+ or collection.get_array() is not None
+ or np.size(collection.get_transforms()) != 0
+ or np.ndim(collection.get_alpha()) != 0
+ ):
+ _native(collection, ax, target, unsupported)
+ return True
+ if isinstance(collection, LineCollection):
+ for i, segment in enumerate(collection.get_segments()):
+ if not len(segment):
+ continue
+ kwargs = _collection_style(collection, i)
+ kwargs.pop("facecolor")
+ kwargs["color"] = kwargs.pop("edgecolor")
+ target.plot(segment[:, 0].copy(), segment[:, 1].copy(), **kwargs)
+ else:
+ for i, path in enumerate(collection.get_paths()):
+ if path.codes is not None and (
+ np.count_nonzero(path.codes == Path.MOVETO) > 1
+ or np.any(
+ ~np.isin(path.codes, [Path.MOVETO, Path.LINETO, Path.CLOSEPOLY])
+ )
+ ):
+ _native(collection, ax, target, unsupported)
+ return True
+ vertices = path.vertices
+ if path.codes is not None and path.codes[-1] == Path.CLOSEPOLY:
+ vertices = vertices[:-1]
+ target.fill(
+ vertices[:, 0].copy(),
+ vertices[:, 1].copy(),
+ hatch=collection.get_hatch(),
+ **_collection_style(collection, i),
+ )
+ return True
+
+
+def _import_axes(ax, target, unsupported):
+ if ax.name != "rectilinear" or type(ax).__module__.startswith("mpl_toolkits."):
+ target._import_projection = ReboundSnapshot(ax, fig=ax.figure)
+ target._import_projection_artist_ids = {
+ name: [id(artist) for artist in getattr(ax, name)]
+ for name in (
+ "lines",
+ "collections",
+ "images",
+ "patches",
+ "texts",
+ "artists",
+ )
+ }
+ unsupported.report.add(
+ ax,
+ "Projection, camera and native artists retained for "
+ "Matplotlib; edit through the rendered axes",
+ fallback="native",
+ )
+ return
+ consumed = set()
+ target.import_groups = {
+ id(container): dict(
+ type=type(container).__name__,
+ label=container.get_label(),
+ artist_ids=[id(artist) for artist in container.get_children()],
+ orientation=getattr(container, "orientation", None),
+ datavalues=copy.deepcopy(getattr(container, "datavalues", None)),
+ )
+ for container in ax.containers
+ }
+ # A bar container retains grouping and orientation, unlike raw rectangles.
+ for container in _tracked(ax.containers, ax, target, unsupported):
+ if isinstance(container, ErrorbarContainer):
+ if _import_errorbar(container, ax, target, unsupported):
+ consumed.update(container.get_children())
+ continue
+ if isinstance(container, StemContainer):
+ unsupported.report.add(
+ container, "Stem components retained as geometry", fallback="geometry"
+ )
+ continue
+ if not isinstance(container, BarContainer):
+ unsupported(f"Unsupported {type(container).__name__}; container skipped")
+ consumed.update(container.get_children())
+ continue
+ for i, patch in enumerate(container.patches):
+ consumed.add(patch)
+ if _needs_native_style(patch):
+ _native(patch, ax, target, unsupported)
+ target.line_data[-1]["source_order"] = ax.get_children().index(patch)
+ continue
+ kwargs = _style(patch)
+ kwargs.update(
+ color=patch.get_facecolor(),
+ edgecolor=patch.get_edgecolor(),
+ linewidth=patch.get_linewidth(),
+ hatch=patch.get_hatch(),
+ linestyle=patch.get_linestyle(),
+ antialiased=patch.get_antialiased(),
+ fill=patch.get_fill(),
+ label=container.get_label() if i == 0 else "",
+ )
+ if container.orientation == "horizontal":
+ target.barh(
+ [patch.get_y() + patch.get_height() / 2],
+ [patch.get_width()],
+ left=patch.get_x(),
+ height=patch.get_height(),
+ **kwargs,
+ )
+ else:
+ target.bar(
+ [patch.get_x() + patch.get_width() / 2],
+ [patch.get_height()],
+ bottom=patch.get_y(),
+ width=patch.get_width(),
+ **kwargs,
+ )
+
+ target.line_data[-1]["source_order"] = ax.get_children().index(patch)
+
+ for line in _tracked(ax.lines, ax, target, unsupported):
+ if line in consumed:
+ continue
+ if (
+ isinstance(line, AxLine)
+ or _needs_native_style(line)
+ or getattr(line._marker, "_user_transform", None) is not None
+ ):
+ _native(line, ax, target, unsupported)
+ continue
+ x, y = line.get_data(orig=False)
+ if (
+ line.get_transform() == ax.get_yaxis_transform()
+ and len(y) == 2
+ and y[0] == y[1]
+ ):
+ target.axhline(y[0], xmin=x[0], xmax=x[1], **_line_style(line))
+ continue
+ if (
+ line.get_transform() == ax.get_xaxis_transform()
+ and len(x) == 2
+ and x[0] == x[1]
+ ):
+ target.axvline(x[0], ymin=y[0], ymax=y[1], **_line_style(line))
+ continue
+ if line.get_transform() != ax.transData:
+ _native(line, ax, target, unsupported)
+ continue
+ target.plot(
+ np.array(x, copy=True),
+ np.array(y, copy=True),
+ **_line_style(line),
+ )
+ for collection in _tracked(ax.collections, ax, target, unsupported):
+ if collection in consumed:
+ continue
+ if _needs_native_style(collection):
+ _native(collection, ax, target, unsupported)
+ continue
+ if _import_collection(collection, ax, target, unsupported):
+ continue
+ if not isinstance(collection, PathCollection):
+ _native(collection, ax, target, unsupported)
+ continue
+ paths = collection.get_paths()
+ if (
+ len(paths) != 1
+ or collection.get_offset_transform() != ax.transData
+ or not isinstance(collection.get_transform(), IdentityTransform)
+ ):
+ _native(collection, ax, target, unsupported)
+ continue
+ offsets = np.ma.asarray(collection.get_offsets()).filled(np.nan)
+ kwargs = _style(collection)
+ kwargs.update(
+ s=collection.get_sizes().copy(),
+ marker=_scatter_marker(paths[0]),
+ edgecolors=collection.get_edgecolors().copy(),
+ linewidths=collection.get_linewidths().copy(),
+ )
+ values = collection.get_array()
+ if values is not None:
+ kwargs.update(
+ c=values.copy(),
+ cmap=copy.copy(collection.get_cmap()),
+ norm=copy.deepcopy(collection.norm),
+ )
+ else:
+ colors = collection.get_facecolors()
+ if len(colors) == 1:
+ kwargs["color"] = tuple(colors[0])
+ elif len(colors):
+ kwargs["c"] = colors.copy()
+ else:
+ kwargs["facecolors"] = "none"
+ target.scatter(offsets[:, 0].copy(), offsets[:, 1].copy(), **kwargs)
+ for image in _tracked(ax.images, ax, target, unsupported):
+ if image.get_transform() != ax.transData or _needs_native_style(image):
+ _native(image, ax, target, unsupported)
+ continue
+ target.add_imshow(
+ image.get_array().copy(),
+ extent=tuple(image.get_extent()),
+ origin=image.origin,
+ cmap=copy.copy(image.get_cmap()),
+ norm=copy.deepcopy(image.norm),
+ interpolation=image.get_interpolation(),
+ interpolation_stage=getattr(image, "_interpolation_stage", "data"),
+ filternorm=image.get_filternorm(),
+ filterrad=image.get_filterrad(),
+ resample=image.get_resample(),
+ **_style(image),
+ )
+ for text in _tracked(ax.texts, ax, target, unsupported):
+ if (
+ text.get_bbox_patch() is not None
+ or _needs_native_style(text)
+ or text.get_wrap()
+ ):
+ _native(text, ax, target, unsupported)
+ continue
+ if isinstance(text, Annotation):
+ if text.xycoords != "data" or text.anncoords != "data":
+ _native(text, ax, target, unsupported)
+ continue
+ arrowprops = text.arrowprops
+ if arrowprops and any(key in arrowprops for key in ("patchA", "patchB")):
+ _native(text, ax, target, unsupported)
+ continue
+ target.annotate(
+ text.get_text(),
+ xy=tuple(text.xy),
+ xytext=tuple(text.get_position()),
+ arrowprops=copy.deepcopy(arrowprops),
+ annotation_clip=text.get_annotation_clip(),
+ **_text_style(text),
+ **_style(text),
+ )
+ continue
+ if text.get_transform() != ax.transData:
+ _native(text, ax, target, unsupported)
+ continue
+ target.text(
+ *text.get_position(),
+ text.get_text(),
+ **_text_style(text),
+ **_style(text),
+ )
+ for patch in _tracked(ax.patches, ax, target, unsupported):
+ if patch in consumed:
+ continue
+ if patch.get_data_transform() != ax.transData or _needs_native_style(patch):
+ _native(patch, ax, target, unsupported)
+ consumed.add(patch)
+ continue
+ kwargs = dict(
+ facecolor=patch.get_facecolor(),
+ edgecolor=patch.get_edgecolor(),
+ linewidth=patch.get_linewidth(),
+ linestyle=patch.get_linestyle(),
+ hatch=patch.get_hatch(),
+ fill=patch.get_fill(),
+ **_style(patch),
+ )
+ if isinstance(patch, StepPatch):
+ data = patch.get_data()
+ target.stairs(
+ data.values.copy(),
+ data.edges.copy(),
+ baseline=copy.deepcopy(data.baseline),
+ orientation=patch.orientation,
+ **kwargs,
+ )
+ consumed.add(patch)
+ elif isinstance(patch, Polygon):
+ xy = patch.get_xy()
+ target.fill(
+ xy[:, 0].copy(), xy[:, 1].copy(), closed=patch.get_closed(), **kwargs
+ )
+ consumed.add(patch)
+ else:
+ _native(patch, ax, target, unsupported)
+ target._import_inset_indicators = {}
+ for artist in _tracked(list(ax.artists) + list(ax.tables), ax, target, unsupported):
+ if isinstance(artist, Legend):
+ continue
+ if (
+ InsetIndicator is not None
+ and isinstance(artist, InsetIndicator)
+ and artist._inset_ax in ax.child_axes
+ ):
+ index = ax.child_axes.index(artist._inset_ax)
+ target._import_inset_indicators[index] = _capture_inset_indicator(artist)
+ unsupported.report.add(
+ artist,
+ "Inset indicator rebuilt from the live inset axes",
+ fallback="native",
+ )
+ continue
+ _native(artist, ax, target, unsupported)
+ target._import_child_axes = []
+ for child in ax.child_axes:
+ _import_child(child, ax, target, unsupported)
+
+ title_kwargs = _text_style(ax.title)
+ title_kwargs["x"] = ax.title.get_position()[0]
+ title_kwargs["pad"] = ax.titleOffsetTrans.transform((0, 0))[1] * 72 / ax.figure.dpi
+ if not ax._autotitlepos:
+ title_kwargs["y"] = ax.title.get_position()[1]
+ target.set_title(ax.get_title(), **title_kwargs)
+ target.set_xlabel(
+ ax.get_xlabel(), labelpad=ax.xaxis.labelpad, **_text_style(ax.xaxis.label)
+ )
+ target.set_ylabel(
+ ax.get_ylabel(), labelpad=ax.yaxis.labelpad, **_text_style(ax.yaxis.label)
+ )
+ target.set_xlim(*ax.get_xlim())
+ target.set_ylim(*ax.get_ylim())
+ for name in ("x", "y"):
+ scale = getattr(ax, f"get_{name}scale")()
+ getattr(target, f"set_{name}scale")(scale)
+ axis = getattr(ax, f"{name}axis")
+ if isinstance(axis.get_major_locator(), FixedLocator):
+ getattr(target, f"set_{name}ticks")(
+ axis.get_majorticklocs().copy(),
+ labels=axis.get_major_formatter().format_ticks(
+ axis.get_majorticklocs()
+ ),
+ )
+ target.set_grid(
+ any(line.get_visible() for line in ax.get_xgridlines() + ax.get_ygridlines())
+ )
+ target.set_legend(ax.get_legend() is not None and ax.get_legend().get_visible())
+ target.set_facecolor(ax.get_facecolor())
+ target.set_axisbelow(ax.get_axisbelow())
+ target.set_aspect(ax.get_aspect())
+ target.set_visible(ax.get_visible())
+ if not ax.axison:
+ target.set_axis_off()
+ for name in (
+ "adjustable",
+ "anchor",
+ "box_aspect",
+ "frame_on",
+ "alpha",
+ "zorder",
+ "rasterized",
+ "autoscalex_on",
+ "autoscaley_on",
+ ):
+ getattr(target, "set_" + name)(getattr(ax, "get_" + name)())
+ target.set_xmargin(ax.margins()[0])
+ target.set_ymargin(ax.margins()[1])
+ unsupported.report.add(
+ ax,
+ "Axis scales, tick and spine styles, and legend layout "
+ "retained for Matplotlib",
+ fallback="native",
+ )
+ target._import_grid = target._grid
+ target._import_scales = (target._xaxis_scale, target._yaxis_scale)
+ target._import_axis_state = capture_axis_state(ax)
+ target._import_legends = [
+ ReboundSnapshot(legend, ax)
+ for legend in ax.get_children()
+ if isinstance(legend, Legend)
+ ]
+ target.line_data.sort(key=lambda entry: entry.get("source_order", 0))
+ for entries in target.layered_line_data.values():
+ entries.sort(key=lambda entry: entry.get("source_order", 0))
+
+
+def _needs_native_style(artist):
+ box = artist.get_clip_box()
+ custom_box = (
+ box is not None
+ and artist.axes is not None
+ and not np.array_equal(box.bounds, artist.axes.bbox.bounds)
+ )
+ return (
+ custom_box
+ or artist.get_path_effects()
+ or artist.get_agg_filter() is not None
+ or artist.get_sketch_params() is not None
+ or artist.get_clip_path() is not None
+ )
+
+
+def _native(artist, ax, target, unsupported):
+ if unsupported.fallback == "skip" or not type(artist).__module__.startswith(
+ ("matplotlib.", "mpl_toolkits.")
+ ):
+ unsupported(f"Unsupported {type(artist).__name__}; artist skipped", artist)
+ return
+ snapshot = ReboundSnapshot(artist, ax)
+ target._add(
+ dict(plot_type="matplotlib_artist", snapshot=snapshot, layer=0, kwargs={}), 0
+ )
+ unsupported.report.add(
+ artist,
+ "Editable artist geometry retained for Matplotlib; "
+ "other backends require a raster import",
+ fallback="native",
+ )
+
+
+def _tracked(artists, ax, target, unsupported):
+ order = {id(artist): i for i, artist in enumerate(ax.get_children())}
+ for artist in artists:
+ unsupported.artist = artist
+ start = len(target.line_data)
+ yield artist
+ entries = target.line_data[start:]
+ children = getattr(artist, "get_children", lambda: [])()
+ index = order.get(
+ id(artist),
+ min(
+ (order.get(id(child), len(order)) for child in children),
+ default=len(order),
+ ),
+ )
+ for entry in entries:
+ entry["source_artist_id"] = id(artist)
+ entry["source_container_ids"] = [
+ identifier
+ for identifier, group in target.import_groups.items()
+ if identifier == id(artist) or id(artist) in group["artist_ids"]
+ ]
+ entry.setdefault("source_order", index)
+ if isinstance(artist, Artist):
+ entry["source_sticky_edges"] = (
+ list(artist.sticky_edges.x),
+ list(artist.sticky_edges.y),
+ )
+ entry["source_picker"] = artist.get_picker()
+ if entries:
+ unsupported.report.add(
+ artist,
+ "Imported " + str(len(entries)) + " plot entries",
+ fallback="geometry" if len(entries) > 1 else None,
+ )
+
+
+def _capture_colorbar(colorbar, target):
+ axis = (
+ colorbar.ax.yaxis if colorbar.orientation == "vertical" else colorbar.ax.xaxis
+ )
+ return dict(
+ target=target,
+ artist_id=id(colorbar.mappable),
+ standalone=(
+ dict(
+ norm=copy.deepcopy(colorbar.mappable.norm),
+ cmap=copy.copy(colorbar.mappable.get_cmap()),
+ )
+ if target is None
+ else None
+ ),
+ position=tuple(colorbar.ax.get_position().bounds),
+ kwargs=dict(
+ orientation=colorbar.orientation,
+ extend=colorbar.extend,
+ extendfrac=colorbar.extendfrac,
+ extendrect=colorbar.extendrect,
+ spacing=colorbar.spacing,
+ drawedges=colorbar.drawedges,
+ boundaries=copy.deepcopy(colorbar.boundaries),
+ values=copy.deepcopy(colorbar.values),
+ ),
+ label=axis.label.get_text(),
+ label_style=_text_style(axis.label),
+ ticks=copy.deepcopy(colorbar.get_ticks()),
+ formatter=ReboundSnapshot(colorbar.formatter, colorbar.ax),
+ state=capture_axis_state(colorbar.ax),
+ )
+
+
+def _capture_inset_indicator(indicator):
+ """Style for ``indicate_inset_zoom``; connectors are recomputed on render.
+
+ ``InsetIndicator`` spans two Axes and recomputes its rectangle/connector
+ geometry from live axes limits on every draw, so it cannot be captured as
+ detached geometry the way single-axes artists are. Recreating it with
+ ``Axes.indicate_inset_zoom`` on the reconstructed parent/inset pair keeps
+ that live behavior instead of freezing a stale snapshot.
+ """
+ rectangle = indicator.rectangle
+ return dict(
+ facecolor=copy.deepcopy(rectangle.get_facecolor()),
+ edgecolor=copy.deepcopy(rectangle.get_edgecolor()),
+ linewidth=rectangle.get_linewidth(),
+ linestyle=rectangle.get_linestyle(),
+ alpha=indicator.get_alpha(),
+ zorder=indicator.get_zorder(),
+ visible=indicator.get_visible(),
+ )
+
+
+def _import_child(child, ax, target, unsupported):
+ from matplotlib.axes._secondary_axes import SecondaryAxis
+
+ from maxplotlib.subfigure.line_plot import LinePlot
+
+ if isinstance(child, SecondaryAxis):
+ target._import_child_axes.append(
+ dict(
+ kind="secondary",
+ orientation=child._orientation,
+ location=child._loc,
+ functions=child._functions,
+ locator=ReboundSnapshot(child.get_axes_locator(), ax),
+ state=capture_axis_state(child),
+ xlabel=child.get_xlabel(),
+ ylabel=child.get_ylabel(),
+ )
+ )
+ else:
+ subplot = LinePlot()
+ _import_axes(child, subplot, unsupported)
+ bounds = ax.transAxes.inverted().transform_bbox(child.bbox).bounds
+ target._import_child_axes.append(
+ dict(
+ kind="inset",
+ bounds=tuple(bounds),
+ subplot=subplot,
+ locator=ReboundSnapshot(child.get_axes_locator(), ax),
+ )
+ )
+ unsupported.report.add(
+ child, "Child axes retained for Matplotlib", fallback="native"
+ )
+
+
+def _raster_import(canvas_cls, figure, axes, whole_figure, report, **kwargs):
+ from matplotlib.backends.backend_agg import FigureCanvasAgg
+
+ snapshot = copy.deepcopy(figure)
+ if not whole_figure:
+ indices = [figure.axes.index(ax) for ax in axes]
+ for i, ax in enumerate(list(snapshot.axes)):
+ if i not in indices:
+ ax.remove()
+ FigureCanvasAgg(snapshot).draw()
+ rgba = np.asarray(snapshot.canvas.buffer_rgba()).copy()
+ kwargs.setdefault("figsize", tuple(figure.get_size_inches()))
+ kwargs.setdefault("dpi", figure.dpi)
+ canvas = canvas_cls(**kwargs)
+ subplot = canvas.add_subplot(row=0, col=0)
+ subplot.add_imshow(rgba)
+ subplot.set_axis_off()
+ subplot.set_grid(False)
+ canvas.import_report = report
+ report.add(
+ figure,
+ "Raster snapshot: data, artist editing, vector paths and interactive "
+ "state are lost",
+ severity="warning",
+ fallback="raster",
+ backends=("matplotlib", "plotly"),
+ )
+ return canvas
diff --git a/src/maxplotlib/backends/tikzfigure/__init__.py b/src/maxplotlib/backends/tikzfigure/__init__.py
new file mode 100644
index 0000000..5f2ddad
--- /dev/null
+++ b/src/maxplotlib/backends/tikzfigure/__init__.py
@@ -0,0 +1,9 @@
+"""The tikzfigure backend: pgfplots figures from what Matplotlib draws.
+
+See :func:`figure_to_tikz`; ``Canvas.render(backend="tikzfigure")`` uses it.
+"""
+
+from .convert import TikzConversionWarning, figure_to_tikz
+from .text import latex
+
+__all__ = ["TikzConversionWarning", "figure_to_tikz", "latex"]
diff --git a/src/maxplotlib/backends/tikzfigure/convert.py b/src/maxplotlib/backends/tikzfigure/convert.py
new file mode 100644
index 0000000..d787777
--- /dev/null
+++ b/src/maxplotlib/backends/tikzfigure/convert.py
@@ -0,0 +1,1331 @@
+"""A drawn Matplotlib figure as a tikzfigure figure of pgfplots axes.
+
+Every axes becomes a pgfplots ``axis`` at the place and size it has in the
+figure, with its limits, scales, labels, title, ticks, grid, spines and
+legend. What is drawn in it is converted artist by artist, in Matplotlib's
+drawing order:
+
+* vector graphics, in pgfplots code: lines and markers
+ (:class:`~matplotlib.lines.Line2D`), scatter plots, line and polygon
+ collections (``hlines``, ``fill_between``, error bars, ...), contour
+ lines, patches (bars, spans, polygons, arrows) and text and annotations;
+* raster images, rendered by Matplotlib itself and placed with
+ ``\\addplot graphics``: meshes, images, filled contours, quivers, and
+ every artist that has no vector counterpart here or would be too large
+ for TeX (``max_markers``, ``max_items``, ``max_points``).
+
+Colorbars are axes of their own: their color strip is an image, their
+ticks and label pgfplots text. Axes that pgfplots cannot represent (polar
+or 3-D projections, symlog scales) become one image each, as do figure
+legends. Text is converted with :func:`~maxplotlib.backends.tikzfigure.text.latex`.
+"""
+
+from __future__ import annotations
+
+import io
+import warnings
+
+import matplotlib.colors as mcolors
+import matplotlib.ticker as mticker
+import numpy as np
+from matplotlib.collections import (
+ Collection,
+ LineCollection,
+ PathCollection,
+ PolyCollection,
+)
+from matplotlib.contour import ContourSet
+from matplotlib.image import AxesImage
+from matplotlib.lines import Line2D
+from matplotlib.markers import MarkerStyle
+from matplotlib.patches import FancyArrowPatch, Patch
+from matplotlib.path import Path
+from matplotlib.text import Annotation, Text
+from matplotlib.transforms import Bbox
+from tikzfigure import TikzFigure
+from tikzfigure.core.axis import Axis2D
+from tikzfigure.core.plot import format_number
+
+from .text import is_multiline, latex
+
+__all__ = ["TikzConversionWarning", "figure_to_tikz"]
+
+
+class TikzConversionWarning(UserWarning):
+ """A part of a Matplotlib figure that is drawn as an image or left out."""
+
+
+# Matplotlib marker -> (pgfplots mark when filled, when hollow, rotation, size factor)
+_MARKS = {
+ "o": ("*", "o", 0, 0.5),
+ ".": ("*", "o", 0, 0.25),
+ ",": ("*", "o", 0, 0.1),
+ "s": ("square*", "square", 0, 0.5),
+ "^": ("triangle*", "triangle", 0, 0.5),
+ "v": ("triangle*", "triangle", 180, 0.5),
+ "<": ("triangle*", "triangle", 90, 0.5),
+ ">": ("triangle*", "triangle", 270, 0.5),
+ "D": ("diamond*", "diamond", 0, 0.5),
+ "d": ("diamond*", "diamond", 0, 0.45),
+ "p": ("pentagon*", "pentagon", 0, 0.5),
+ "h": ("pentagon*", "pentagon", 0, 0.5),
+ "H": ("pentagon*", "pentagon", 0, 0.5),
+ "*": ("star", "star", 0, 0.5),
+ "+": ("+", "+", 0, 0.5),
+ "P": ("+", "+", 0, 0.5),
+ "x": ("x", "x", 0, 0.5),
+ "X": ("x", "x", 0, 0.5),
+ "|": ("|", "|", 0, 0.5),
+ "_": ("-", "-", 0, 0.5),
+ "1": ("Mercedes star", "Mercedes star", 180, 0.5),
+ "2": ("Mercedes star", "Mercedes star", 0, 0.5),
+}
+_NO_MARKER = ("None", "none", "", " ", None)
+_NO_LINE = ("None", "none", "", " ")
+_HA = {"left": "west", "center": "", "right": "east"}
+_VA = {
+ "top": "north",
+ "center": "",
+ "bottom": "south",
+ "baseline": "base",
+ "center_baseline": "mid",
+}
+# formatters whose labels pgfplots writes just as well by itself
+_AUTO_FORMATTERS = (
+ mticker.ScalarFormatter,
+ mticker.LogFormatterSciNotation,
+ mticker.LogFormatterMathtext,
+)
+_AUTO_LOCATORS = (
+ mticker.AutoLocator,
+ mticker.MaxNLocator,
+ mticker.LogLocator,
+ mticker.AutoMinorLocator,
+)
+
+
+def figure_to_tikz(
+ figure,
+ *,
+ raster_dpi: float = 300,
+ max_markers: int = 2000,
+ max_items: int = 500,
+ max_points: int = 20000,
+ precision: int = 6,
+) -> TikzFigure:
+ """Convert a Matplotlib figure into a :class:`~tikzfigure.TikzFigure`.
+
+ The figure is drawn first (without showing it), so that its layout,
+ autoscaled limits, ticks and legend positions are final, and is left
+ as it was afterwards.
+
+ Parameters
+ ----------
+ figure : matplotlib.figure.Figure
+ The figure to convert.
+ raster_dpi : float, optional
+ Resolution of the parts drawn as images (meshes, images, ...).
+ Default: 300.
+ max_markers : int, optional
+ Scatter plots with more points are drawn as an image. Default: 2000.
+ max_items : int, optional
+ Collections with more differently styled items (e.g. a line
+ collection colored by value) are drawn as an image. Default: 500.
+ max_points : int, optional
+ Lines with more points, after Matplotlib's path simplification, are
+ drawn as an image. Default: 20000.
+ precision : int, optional
+ Significant digits of the written coordinates. Default: 6.
+
+ Returns
+ -------
+ tikzfigure.TikzFigure
+ The figure: ``.generate_tikz()`` for the code, ``.savefig("f.pdf")``
+ to compile it, ``.savefig("f.tikz")`` to write the code and its
+ images.
+ """
+ converter = _FigureConverter(
+ figure,
+ raster_dpi=raster_dpi,
+ max_markers=max_markers,
+ max_items=max_items,
+ max_points=max_points,
+ precision=precision,
+ )
+ return converter.convert()
+
+
+class _Colors:
+ """Colors used in the figure, defined once as ``\\definecolor``."""
+
+ def __init__(self):
+ self.names: dict[str, str] = {}
+
+ def __call__(self, color):
+ """``(name, opacity)`` of a color; name ``None`` for a transparent one."""
+ rgba = mcolors.to_rgba(color)
+ if rgba[3] <= 0:
+ return None, 0.0
+ code = mcolors.to_hex(rgba, keep_alpha=False)[1:].upper()
+ name = self.names.setdefault(code, f"mpl{code}")
+ return name, float(rgba[3])
+
+ def definitions(self) -> str:
+ return "\n".join(
+ f"\\definecolor{{{name}}}{{HTML}}{{{code}}}"
+ for code, name in self.names.items()
+ )
+
+
+def _font(size) -> str:
+ size = float(size)
+ return f"\\fontsize{{{size:g}}}{{{1.2 * size:g}}}\\selectfont"
+
+
+def _pt(value) -> str:
+ return f"{float(value):.4g}pt"
+
+
+def _dash(pattern) -> str | None:
+ """A Matplotlib ``(offset, [on, off, ...])`` dash pattern in points as TikZ."""
+ if pattern is None:
+ return None
+ offset, sequence = pattern
+ if not sequence:
+ return None
+ parts = []
+ for index, length in enumerate(sequence):
+ parts.append(("on" if index % 2 == 0 else "off") + " " + _pt(length))
+ option = "dash pattern=" + " ".join(parts)
+ if offset:
+ option += f", dash phase={_pt(offset)}"
+ return option
+
+
+class _FigureConverter:
+ def __init__(self, figure, **options):
+ self.figure = figure
+ self.options = options
+ self.colors = _Colors()
+ self.tikz = TikzFigure()
+ self.width, self.height = figure.get_size_inches()
+
+ # -- the figure ----------------------------------------------------------
+ def convert(self) -> TikzFigure:
+ figure = self.figure
+ figure.draw_without_rendering()
+ # images are rendered from this figure; a layout engine would move the
+ # axes while other parts are hidden
+ engine = figure.get_layout_engine()
+ figure.set_layout_engine("none")
+ try:
+ for ax in figure.axes:
+ self._axes_tree(ax)
+ for text in figure.texts:
+ self._figure_text(text)
+ for legend in figure.legends:
+ if legend.get_visible():
+ self._as_image(legend, "a figure legend")
+ finally:
+ figure._layout_engine = engine
+ self.tikz.add_package("lmodern")
+ self.tikz.add_package("amsmath")
+ definitions = self.colors.definitions()
+ if definitions:
+ self.tikz.add_raw(definitions)
+ return self.tikz
+
+ def _axes_tree(self, ax):
+ if not ax.get_visible():
+ return
+ if ax.name != "rectilinear" or not _supported_scales(ax):
+ what = (
+ f"{ax.name} axes"
+ if ax.name != "rectilinear"
+ else f"{ax.get_xscale()}/{ax.get_yscale()} scaled axes"
+ )
+ self._as_image(ax, what)
+ return
+ _AxesConverter(self, ax).convert()
+ for child in ax.child_axes:
+ self._axes_tree(child)
+
+ def _figure_text(self, text):
+ if not text.get_visible() or not text.get_text().strip():
+ return
+ if isinstance(text, Annotation):
+ text.update_positions(self.figure._get_renderer())
+ x, y = _text_display_position(text)
+ if not self.figure.bbox.contains(x, y):
+ return
+ x_in, y_in = x / self.figure.dpi, y / self.figure.dpi
+ node = _text_node(text, self.colors, f"({x_in:.4f}in,{y_in:.4f}in)")
+ self.tikz.add_raw(node)
+
+ # -- images of whole parts ---------------------------------------------
+ def _as_image(self, artist, what):
+ """Draw ``artist`` (axes or legend) with its decorations as one image."""
+ warnings.warn(
+ f"{what} cannot be drawn with pgfplots; it is included as an image",
+ TikzConversionWarning,
+ stacklevel=4,
+ )
+ renderer = self.figure._get_renderer()
+ bbox = artist.get_tightbbox(renderer)
+ if bbox is None or bbox.width <= 0 or bbox.height <= 0:
+ return
+ inches = bbox.transformed(self.figure.dpi_scale_trans.inverted())
+ data = self.render_only([artist], inches)
+ axis = Axis2D(
+ xlim=(0, 1),
+ ylim=(0, 1),
+ grid=False,
+ width=f"{inches.width:.4f}in",
+ height=f"{inches.height:.4f}in",
+ options=[
+ "hide axis",
+ "scale only axis",
+ "anchor=south west",
+ f"at={{({inches.x0:.4f}in,{inches.y0:.4f}in)}}",
+ ],
+ )
+ axis.add_graphics(0, 1, 0, 1, data=data, plot_options=["forget plot"])
+ self.tikz.axes.append(axis)
+
+ def render_only(self, artists, inches: Bbox, keep_axes=None) -> bytes:
+ """A PNG of the figure region ``inches`` showing only ``artists``.
+
+ ``keep_axes`` is the axes whose decorations stay hidden while its
+ artists in ``artists`` are drawn; every other axes is hidden.
+ """
+ figure = self.figure
+ keep = set(map(id, artists))
+ hidden = []
+
+ def hide(artist):
+ if id(artist) not in keep and artist.get_visible():
+ artist.set_visible(False)
+ hidden.append(artist)
+
+ # the axes containing a kept artist, and the axes containing those,
+ # stay visible (an artist of a hidden axes is not drawn), with their
+ # other children hidden
+ parents = {id(child): ax for ax in _all_axes(figure) for child in ax.child_axes}
+ owners = set()
+ for artist in artists:
+ owner = getattr(artist, "axes", None)
+ if owner is artist:
+ owner = parents.get(id(artist))
+ while owner is not None:
+ owners.add(id(owner))
+ owner = parents.get(id(owner))
+ for ax in _all_axes(figure):
+ if id(ax) in keep:
+ continue
+ if id(ax) in owners:
+ for child in ax.get_children():
+ if id(child) not in owners:
+ hide(child)
+ else:
+ hide(ax)
+ hide(figure.patch)
+ for text in figure.texts:
+ hide(text)
+ for legend in figure.legends:
+ hide(legend)
+ for artist in artists: # the artists themselves are drawn
+ if not artist.get_visible():
+ artist.set_visible(True)
+ hidden.append(("shown", artist))
+ buffer = io.BytesIO()
+ try:
+ figure.savefig(
+ buffer,
+ format="png",
+ dpi=self.options["raster_dpi"],
+ transparent=True,
+ bbox_inches=inches,
+ pad_inches=0,
+ )
+ finally:
+ for artist in hidden:
+ if isinstance(artist, tuple):
+ artist[1].set_visible(False)
+ else:
+ artist.set_visible(True)
+ return buffer.getvalue()
+
+
+def _all_axes(figure):
+ out = []
+
+ def walk(ax):
+ out.append(ax)
+ for child in ax.child_axes:
+ walk(child)
+
+ for ax in figure.axes:
+ walk(ax)
+ return out
+
+
+def _supported_scales(ax) -> bool:
+ return ax.get_xscale() in ("linear", "log") and ax.get_yscale() in (
+ "linear",
+ "log",
+ )
+
+
+def _text_display_position(text):
+ return text.get_transform().transform(text.get_unitless_position())
+
+
+def _anchor(text) -> str:
+ vertical = _VA.get(text.get_va(), "")
+ horizontal = _HA.get(text.get_ha(), "")
+ if vertical == "base" and not horizontal:
+ return "base"
+ if vertical == "mid" and not horizontal:
+ return "mid"
+ anchor = " ".join(part for part in (vertical, horizontal) if part)
+ return anchor or "center"
+
+
+def _text_node(text, colors, at: str) -> str:
+ """A ``\\node`` for a Matplotlib text, placed at the TikZ point ``at``."""
+ options = [f"anchor={_anchor(text)}", "inner sep=0pt"]
+ rotation = text.get_rotation()
+ if rotation:
+ options.append(f"rotate={rotation:g}")
+ color, opacity = colors(text.get_color())
+ if color is not None and color != "mpl000000":
+ options.append(f"text={color}")
+ alpha = text.get_alpha()
+ if alpha is not None and alpha < 1:
+ opacity *= alpha
+ if opacity < 1:
+ options.append(f"text opacity={opacity:.3g}")
+ options.append(f"font={{{_font(text.get_fontsize())}}}")
+ if is_multiline(text.get_text()):
+ alignment = getattr(text, "_multialignment", None) or text.get_ha()
+ options.append(f"align={alignment}")
+ box = text.get_bbox_patch()
+ if box is not None:
+ face, face_opacity = colors(box.get_facecolor())
+ edge, _ = colors(box.get_edgecolor())
+ options[1] = "inner sep=2pt"
+ if face is not None:
+ options.append(f"fill={face}")
+ if face_opacity < 1:
+ options.append(f"fill opacity={face_opacity:.3g}, text opacity=1")
+ if edge is not None and box.get_linewidth() > 0:
+ options.append(f"draw={edge}, line width={_pt(box.get_linewidth())}")
+ if "round" in type(box.get_boxstyle()).__name__.lower():
+ options.append("rounded corners=2pt")
+ return f"\\node[{', '.join(options)}] at {at} {{{latex(text.get_text())}}};"
+
+
+class _AxesConverter:
+ """One Matplotlib axes as one pgfplots axis."""
+
+ def __init__(self, parent: _FigureConverter, ax):
+ self.parent = parent
+ self.figure = parent.figure
+ self.ax = ax
+ self.colors = parent.colors
+ self.options = parent.options
+ self.colorbar = getattr(ax, "_colorbar", None)
+ self.legend_entries: list = []
+ self.has_images = False
+
+ # -- the axis ----------------------------------------------------------
+ def convert(self):
+ ax = self.ax
+ position = ax.get_position()
+ width, height = self.parent.width, self.parent.height
+ x0, y0 = position.x0 * width, position.y0 * height
+ self.size = (position.width * width, position.height * height)
+ self.inches = Bbox.from_bounds(x0, y0, *self.size)
+ xlim, ylim = ax.get_xlim(), ax.get_ylim()
+
+ options = [
+ "scale only axis",
+ "anchor=south west",
+ f"at={{({x0:.4f}in,{y0:.4f}in)}}",
+ "clip mode=individual",
+ "unbounded coords=jump",
+ "every axis plot/.append style={line join=round}",
+ ]
+ if xlim[0] > xlim[1]:
+ options.append("x dir=reverse")
+ if ylim[0] > ylim[1]:
+ options.append("y dir=reverse")
+ if not ax.axison:
+ options.append("hide axis")
+ else:
+ options.extend(self._frame())
+ for name in ("x", "y"):
+ options.extend(self._ticks(name))
+ options.extend(self._labels())
+ options.extend(self._background())
+
+ self.axis = Axis2D(
+ xlabel=latex(ax.get_xlabel()) if ax.xaxis.get_visible() else "",
+ ylabel=latex(ax.get_ylabel()) if ax.yaxis.get_visible() else "",
+ title=latex(self._title()),
+ xlim=(float(min(xlim)), float(max(xlim))),
+ ylim=(float(min(ylim)), float(max(ylim))),
+ xlog=ax.get_xscale() == "log",
+ ylog=ax.get_yscale() == "log",
+ grid=None,
+ width=f"{self.size[0]:.4f}in",
+ height=f"{self.size[1]:.4f}in",
+ options=options,
+ )
+ self._legend_setup()
+ self._contents()
+ self._legend_entries()
+ if self.has_images:
+ self.axis.options.append("axis on top")
+ grid = self._grid()
+ if grid:
+ self.axis.options.extend(grid)
+ self.parent.tikz.axes.append(self.axis)
+
+ def _title(self) -> str:
+ for location in ("center", "left", "right"):
+ title = self.ax.get_title(location)
+ if title:
+ return title
+ return ""
+
+ def _frame(self) -> list[str]:
+ spines = self.ax.spines
+ visible = {
+ name: name in spines and spines[name].get_visible()
+ for name in ("left", "right", "top", "bottom")
+ }
+ if "outline" in spines and spines["outline"].get_visible(): # a colorbar
+ return []
+ if all(visible.values()):
+ return []
+ options = []
+ x = (
+ "box"
+ if visible["bottom"] and visible["top"]
+ else "bottom" if visible["bottom"] else "top" if visible["top"] else None
+ )
+ y = (
+ "box"
+ if visible["left"] and visible["right"]
+ else "left" if visible["left"] else "right" if visible["right"] else None
+ )
+ options.append(f"axis x line*={x}" if x else "axis x line=none")
+ options.append(f"axis y line*={y}" if y else "axis y line=none")
+ return options
+
+ def _ticks(self, name) -> list[str]:
+ ax = self.ax
+ axis = getattr(ax, f"{name}axis")
+ if not axis.get_visible():
+ return [f"{name}tick=\\empty", f"{name}ticklabels={{}}"]
+ options = []
+ locator, formatter = axis.get_major_locator(), axis.get_major_formatter()
+ lo, hi = sorted(getattr(ax, f"get_{name}lim")())
+ # the ticks where Matplotlib has them; their labels written by pgfplots,
+ # unless Matplotlib's formatter writes something else than numbers
+ tolerance = 1e-9 * abs(hi - lo)
+ locs = [
+ loc
+ for loc in axis.get_majorticklocs()
+ if lo - tolerance <= loc <= hi + tolerance
+ ]
+ if isinstance(locator, mticker.NullLocator) or not locs:
+ options.append(f"{name}tick=\\empty")
+ else:
+ options.append(
+ f"{name}tick={{{','.join(f'{float(loc):.10g}' for loc in locs)}}}"
+ )
+ if not isinstance(formatter, _AUTO_FORMATTERS):
+ labels = formatter.format_ticks(locs)
+ options.append(
+ f"{name}ticklabels={{{','.join('{' + latex(lab) + '}' for lab in labels)}}}"
+ )
+ ticks = axis.get_major_ticks()
+ if ticks:
+ tick = ticks[0]
+ first = tick.tick1line.get_visible()
+ second = tick.tick2line.get_visible()
+ side = "both" if first and second else "left" if first else "right"
+ if not first and not second:
+ options.append(f"major {name} tick style={{draw=none}}")
+ else:
+ options.append(f"{name}tick pos={side}")
+ label1, label2 = tick.label1.get_visible(), tick.label2.get_visible()
+ if not label1 and not label2:
+ options.append(f"{name}ticklabels={{}}")
+ options.append(f"scaled {name} ticks=false")
+ else:
+ options.append(
+ f"{name}ticklabel pos={'right' if label2 and not label1 else 'left'}"
+ )
+ options.append(
+ f"{name} tick label style={{font={{{_font(tick.label1.get_fontsize())}}}}}"
+ )
+ if name == "x":
+ direction = getattr(tick, "_tickdir", "out")
+ align = {"in": "inside", "out": "outside", "inout": "center"}.get(
+ direction, "outside"
+ )
+ options.append(f"tick align={align}")
+ options.append(f"major tick length={_pt(tick._size)}")
+ if axis.get_label_position() in ("top", "right"):
+ options.append(f"{name}label near ticks")
+ if not any(option.startswith(f"{name}ticklabel pos") for option in options):
+ options.append(f"{name}ticklabel pos=right")
+ return options
+
+ def _labels(self) -> list[str]:
+ options = []
+ for key, label in (
+ ("xlabel", self.ax.xaxis.label),
+ ("ylabel", self.ax.yaxis.label),
+ ("title", self.ax.title),
+ ):
+ style = [f"font={{{_font(label.get_fontsize())}}}"]
+ color, _ = self.colors(label.get_color())
+ if color is not None and color != "mpl000000":
+ style.append(f"text={color}")
+ if is_multiline(label.get_text()):
+ style.append("align=center")
+ options.append(f"{key} style={{{', '.join(style)}}}")
+ return options
+
+ def _background(self) -> list[str]:
+ if self.colorbar is not None:
+ return []
+ patch = self.ax.patch
+ if not patch.get_visible():
+ return []
+ color, opacity = self.colors(patch.get_facecolor())
+ if color is None or color == "mplFFFFFF":
+ return []
+ fill = f"fill={color}" + (
+ f", fill opacity={opacity:.3g}" if opacity < 1 else ""
+ )
+ return [f"axis background/.style={{{fill}}}"]
+
+ def _grid(self) -> list[str]:
+ options = []
+ style = None
+ for name in ("x", "y"):
+ axis = getattr(self.ax, f"{name}axis")
+ ticks = axis.get_major_ticks()
+ if ticks and ticks[0].gridline.get_visible():
+ options.append(f"{name}majorgrids")
+ style = style or ticks[0].gridline
+ if style is not None:
+ color, opacity = self.colors(style.get_color())
+ alpha = style.get_alpha()
+ if alpha is not None:
+ opacity *= alpha
+ parts = [f"draw={color}", f"line width={_pt(style.get_linewidth())}"]
+ if opacity < 1:
+ parts.append(f"draw opacity={opacity:.3g}")
+ dash = _dash(style._dash_pattern) if style.is_dashed() else None
+ parts.append(dash or "solid")
+ options.append(f"major grid style={{{', '.join(parts)}}}")
+ return options
+
+ # -- the legend --------------------------------------------------------
+ def _legend_setup(self):
+ legend = self.ax.get_legend()
+ if legend is None or not legend.get_visible():
+ return
+ handles = getattr(legend, "legend_handles", None)
+ if handles is None: # Matplotlib < 3.7
+ handles = legend.legendHandles
+ self.legend_entries = [
+ (handle, text.get_text())
+ for handle, text in zip(handles, legend.get_texts())
+ ]
+ bbox = legend.get_window_extent().transformed(self.ax.transAxes.inverted())
+ center_x, center_y = (bbox.x0 + bbox.x1) / 2, (bbox.y0 + bbox.y1) / 2
+ vertical = "north" if center_y > 0.5 else "south"
+ horizontal = "east" if center_x > 0.5 else "west"
+ at = (
+ bbox.x1 if horizontal == "east" else bbox.x0,
+ bbox.y1 if vertical == "north" else bbox.y0,
+ )
+ style = []
+ frame = legend.get_frame()
+ if legend.get_frame_on() and frame.get_visible():
+ face, face_opacity = self.colors(frame.get_facecolor())
+ edge, _ = self.colors(frame.get_edgecolor())
+ style.append(f"fill={face}" if face else "fill=none")
+ if face_opacity < 1:
+ style.append(f"fill opacity={face_opacity:.3g}, text opacity=1")
+ style.append(f"draw={edge}" if edge else "draw=none")
+ if "round" in type(frame.get_boxstyle()).__name__.lower():
+ style.append("rounded corners=2pt")
+ else:
+ style.extend(["draw=none", "fill=none"])
+ if legend.get_texts():
+ style.append(f"font={{{_font(legend.get_texts()[0].get_fontsize())}}}")
+ style.append("cells={anchor=west}")
+ self.axis.set_legend(
+ at=at,
+ anchor=f"{vertical} {horizontal}",
+ columns=getattr(legend, "_ncols", 1) or None,
+ style=style,
+ )
+
+ def _legend_label(self, artist) -> str:
+ """Plots never make legend entries themselves; see :meth:`_legend_entries`."""
+ return ""
+
+ def _legend_entries(self):
+ """The legend as Matplotlib draws it: an image and a text per entry.
+
+ The plots are left out of the legend (``forget plot``), so that the
+ entries keep the legend's order and style, also for artists drawn as
+ images.
+ """
+ for handle, label in self.legend_entries:
+ options = self._legend_image(handle)
+ self.axis.add_raw(
+ f"\\addlegendimage{{{', '.join(options)}}}\n"
+ f"\\addlegendentry{{{latex(label)}}}"
+ )
+
+ def _legend_image(self, handle) -> list[str]:
+ if isinstance(handle, Line2D):
+ options = self._line_options(handle)
+ if options:
+ return options
+ elif isinstance(handle, Patch):
+ face = handle.get_facecolor() if handle.get_fill() else "none"
+ # the patch's colors include its alpha
+ options = self._area_options(
+ face,
+ handle.get_edgecolor(),
+ handle.get_linewidth(),
+ _dash(handle._dash_pattern),
+ )
+ if options:
+ return options
+ elif isinstance(handle, PathCollection):
+ handle.update_scalarmappable()
+ faces, edges = handle.get_facecolors(), handle.get_edgecolors()
+ sizes = handle.get_sizes()
+ widths = np.atleast_1d(handle.get_linewidths())
+ paths = handle.get_paths()
+ return ["only marks"] + self._mark(
+ _marker_name(paths[0]) if len(paths) else "o",
+ float(np.sqrt(sizes[0])) if len(sizes) else 6.0,
+ faces[0] if len(faces) else "none",
+ edges[0] if len(edges) and not isinstance(edges, str) else "none",
+ float(widths[0]) if len(widths) else 1.0,
+ )
+ elif isinstance(handle, LineCollection):
+ colors = handle.get_colors()
+ widths = np.atleast_1d(handle.get_linewidths())
+ styles = handle.get_linestyles()
+ stroke = self._stroke(
+ colors[0] if len(colors) else "black",
+ float(widths[0]) if len(widths) else 1.0,
+ _dash(styles[0]) if len(styles) else None,
+ )
+ if stroke:
+ return stroke + ["mark=none"]
+ return ["empty legend"]
+
+ # -- what is drawn -----------------------------------------------------
+ def _contents(self):
+ ax = self.ax
+ children = list(getattr(ax, "_children", []))
+ if not children: # an older Matplotlib
+ children = (
+ ax.collections
+ + ax.patches
+ + ax.lines
+ + ax.texts
+ + ax.images
+ + ax.tables
+ )
+ children = [child for child in children if child.get_visible()]
+ children.sort(key=lambda artist: artist.get_zorder()) # stable: drawing order
+ pending_images = []
+ for artist in children:
+ emit = None if self.colorbar is not None else self._vector(artist)
+ if emit is None:
+ pending_images.append(artist)
+ continue
+ self._flush_images(pending_images)
+ pending_images = []
+ emit()
+ self._flush_images(pending_images)
+
+ def _flush_images(self, artists):
+ if not artists:
+ return
+ reasons = sorted(
+ {
+ type(artist).__name__
+ for artist in artists
+ if not isinstance(artist, _ALWAYS_RASTER)
+ }
+ )
+ if reasons and self.colorbar is None:
+ warnings.warn(
+ f"drawing {', '.join(reasons)} as an image in the pgfplots axis",
+ TikzConversionWarning,
+ stacklevel=6,
+ )
+ data = self.parent.render_only(artists, self.inches)
+ xlim, ylim = self.ax.get_xlim(), self.ax.get_ylim()
+ self.axis.add_graphics(
+ float(min(xlim)),
+ float(max(xlim)),
+ float(min(ylim)),
+ float(max(ylim)),
+ data=data,
+ plot_options=["forget plot"],
+ )
+ self.has_images = True
+
+ def _vector(self, artist):
+ """A function adding ``artist`` to the axis as vector graphics, or None."""
+ if isinstance(artist, Line2D):
+ return self._line(artist)
+ if isinstance(artist, ContourSet):
+ return self._contour_lines(artist)
+ if isinstance(artist, PathCollection):
+ return self._scatter(artist)
+ if isinstance(artist, LineCollection):
+ return self._line_collection(artist)
+ if type(artist) is PolyCollection or type(artist).__name__ in (
+ "FillBetweenPolyCollection",
+ ):
+ return self._polygons(artist)
+ if isinstance(artist, Patch):
+ return self._patch(artist)
+ if isinstance(artist, Text):
+ return self._text(artist)
+ return None
+
+ # coordinates
+ def _view(self):
+ """The axes box in display coordinates, grown by one box size on each side.
+
+ Geometry is clipped to it: what lies further out is not shown, and its
+ coordinates could exceed the largest dimension TeX can hold.
+ """
+ box = self.ax.bbox
+ return (
+ box.x0 - box.width,
+ box.y0 - box.height,
+ box.x1 + box.width,
+ box.y1 + box.height,
+ )
+
+ def _inside(self, display):
+ x0, y0, x1, y1 = self._view()
+ display = np.asarray(display, dtype=float).reshape(-1, 2)
+ with np.errstate(invalid="ignore"):
+ return (
+ (display[:, 0] >= x0)
+ & (display[:, 0] <= x1)
+ & (display[:, 1] >= y0)
+ & (display[:, 1] <= y1)
+ )
+
+ def _clip_line(self, display, simplify=False):
+ """A polyline in display coordinates clipped to :meth:`_view`; gaps as nan.
+
+ With ``simplify``, it is also simplified as Matplotlib does when drawing.
+ """
+ display = np.asarray(display, dtype=float).reshape(-1, 2)
+ if not simplify and self._inside(display).all():
+ return display
+ finite = np.isfinite(display).all(axis=1)
+ if not finite.any():
+ return np.empty((0, 2))
+ path = Path(display).cleaned(
+ remove_nans=True, clip=self._view(), simplify=simplify
+ )
+ return _with_gaps(path)
+
+ def _clip_polygon(self, display):
+ """A polygon in display coordinates clipped to :meth:`_view` (empty if outside)."""
+ display = np.asarray(display, dtype=float).reshape(-1, 2)
+ display = display[np.isfinite(display).all(axis=1)]
+ if len(display) < 3 or self._inside(display).all():
+ return display
+ closed = Path(np.vstack([display, display[:1]]), closed=True)
+ clipped = closed.clip_to_bbox(Bbox.from_extents(*self._view()))
+ polygons = clipped.to_polygons()
+ return polygons[0] if polygons else np.empty((0, 2))
+
+ def _to_data(self, display):
+ display = np.asarray(display, dtype=float).reshape(-1, 2)
+ with np.errstate(all="ignore"):
+ data = self.ax.transData.inverted().transform(display)
+ # display -> data leaves round-off (1e-16 for 0): snap it on linear axes
+ for index, name in enumerate(("x", "y")):
+ if getattr(self.ax, f"get_{name}scale")() == "linear":
+ lo, hi = getattr(self.ax, f"get_{name}lim")()
+ tiny = 1e-9 * abs(hi - lo)
+ with np.errstate(invalid="ignore"):
+ data[np.abs(data[:, index]) < tiny, index] = 0.0
+ return data
+
+ def _add(self, data, options, label="", cycle=False, clip=True):
+ data = np.asarray(data, dtype=float)
+ if not clip:
+ self.axis.add_raw(self._draw(data, options, cycle))
+ return
+ self.axis.add_plot(
+ x=data[:, 0].tolist(),
+ y=data[:, 1].tolist(),
+ label=label,
+ options=list(options) + ["forget plot"],
+ cycle=cycle,
+ precision=self.options["precision"],
+ )
+
+ def _draw(self, data, options, cycle) -> str:
+ """A ``\\draw`` path in axis coordinates, which pgfplots does not clip."""
+ precision = self.options["precision"]
+ keep = [
+ option
+ for option in options
+ if not option.startswith(
+ ("mark", "only marks", "forget plot", "area legend")
+ )
+ ]
+ parts = []
+ connect = False
+ for x, y in data:
+ if not (np.isfinite(x) and np.isfinite(y)):
+ connect = False
+ continue
+ point = f"(axis cs:{format_number(float(x), precision)},{format_number(float(y), precision)})"
+ parts.append(("-- " if connect else "") + point)
+ connect = True
+ if cycle and parts:
+ parts.append("-- cycle")
+ return f"\\draw[{', '.join(keep)}] {' '.join(parts)};"
+
+ def _stroke(self, color, linewidth, dash, alpha=None) -> list[str] | None:
+ name, opacity = self.colors(color)
+ if name is None or linewidth <= 0:
+ return None
+ if alpha is not None:
+ opacity *= alpha
+ options = [f"draw={name}", f"line width={_pt(linewidth)}"]
+ if opacity < 1:
+ options.append(f"draw opacity={opacity:.3g}")
+ options.append(dash or "solid")
+ return options
+
+ def _fill(self, color, alpha=None) -> list[str] | None:
+ name, opacity = self.colors(color)
+ if name is None:
+ return None
+ if alpha is not None:
+ opacity *= alpha
+ options = [f"fill={name}"]
+ if opacity < 1:
+ options.append(f"fill opacity={opacity:.3g}")
+ return options
+
+ def _mark(self, marker, size, face, edge, edge_width, alpha=None) -> list[str]:
+ """``mark=...`` options for a Matplotlib marker of ``size`` points."""
+ if isinstance(marker, str) and marker in _MARKS:
+ filled_mark, hollow_mark, rotation, factor = _MARKS[marker]
+ else:
+ filled_mark, hollow_mark, rotation, factor = _MARKS["o"]
+ fill = self._fill(face, alpha)
+ mark = filled_mark if fill else hollow_mark
+ mark_options = ["solid"]
+ if rotation:
+ mark_options.append(f"rotate={rotation}")
+ mark_options.extend(fill or ["fill=none"])
+ stroke = self._stroke(edge, edge_width, None, alpha)
+ if stroke:
+ mark_options.extend(option for option in stroke if option != "solid")
+ else:
+ mark_options.append("draw=none")
+ return [
+ f"mark={mark}",
+ f"mark size={_pt(max(size * factor, 0.1))}",
+ f"mark options={{{', '.join(mark_options)}}}",
+ ]
+
+ # artists
+ def _line(self, line):
+ path = line.get_path()
+ if not len(path.vertices):
+ return lambda: None
+ display = line.get_transform().transform(path.vertices)
+ marker = line.get_marker()
+ has_marker = marker not in _NO_MARKER and line.get_markersize() > 0
+ if has_marker: # the markers stay where they are; far ones are left out
+ display = np.where(self._inside(display)[:, None], display, np.nan)
+ else:
+ display = self._clip_line(display, simplify=len(display) > 2000)
+ if len(display) > self.options["max_points"]:
+ return None
+ data = self._to_data(display)
+ options = self._line_options(line)
+ if options is None:
+ return lambda: None
+ label = self._legend_label(line)
+ clip = line.get_clip_on()
+ return lambda: self._add(data, options, label, clip=clip)
+
+ def _line_options(self, line) -> list[str] | None:
+ """The style of a Line2D as plot options; None if it draws nothing."""
+ marker = line.get_marker()
+ has_marker = marker not in _NO_MARKER and line.get_markersize() > 0
+ has_line = line.get_linestyle() not in _NO_LINE and line.get_linewidth() > 0
+ alpha = line.get_alpha()
+ options = []
+ if has_line:
+ dash = _dash(line._dash_pattern) if line.is_dashed() else None
+ stroke = self._stroke(line.get_color(), line.get_linewidth(), dash, alpha)
+ if stroke is None:
+ has_line = False
+ else:
+ options.extend(stroke)
+ if not has_line:
+ if not has_marker:
+ return None
+ options.append("only marks")
+ if has_marker:
+ face = line.get_markerfacecolor()
+ if line.get_fillstyle() == "none":
+ face = "none"
+ options.extend(
+ self._mark(
+ marker,
+ line.get_markersize(),
+ face,
+ line.get_markeredgecolor(),
+ line.get_markeredgewidth(),
+ alpha,
+ )
+ )
+ else:
+ options.append("mark=none")
+ return options
+
+ def _contour_lines(self, contours):
+ if contours.filled:
+ return None
+ paths = contours.get_paths()
+ transform = contours.get_transform()
+ colors = contours.get_edgecolor()
+ widths = np.atleast_1d(contours.get_linewidth())
+ styles = contours.get_linestyle()
+ pieces = []
+ total = 0
+ for index, path in enumerate(paths):
+ lines = transform.transform_path(path).to_polygons(closed_only=False)
+ if not lines:
+ continue
+ joined = []
+ for line in lines:
+ joined.append(self._clip_line(line))
+ joined.append([[np.nan, np.nan]])
+ display = np.concatenate(joined[:-1])
+ total += len(display)
+ dash = _dash(styles[index % len(styles)]) if len(styles) else None
+ stroke = self._stroke(
+ colors[index % len(colors)], widths[index % len(widths)], dash
+ )
+ if stroke is not None:
+ pieces.append((self._to_data(display), stroke + ["mark=none"]))
+ if total > self.options["max_points"]:
+ return None
+ texts = list(getattr(contours, "labelTexts", []))
+
+ def emit():
+ for data, options in pieces:
+ self._add(data, options)
+ for text in texts:
+ self._text(text)()
+
+ return emit
+
+ def _scatter(self, collection):
+ offsets = np.ma.asarray(collection.get_offsets()).filled(np.nan)
+ paths = collection.get_paths()
+ count = len(offsets)
+ if count > self.options["max_markers"] or len(paths) != 1:
+ return None
+ if count == 0:
+ return lambda: None
+ display = collection.get_offset_transform().transform(offsets)
+ inside = self._inside(display)
+ data = self._to_data(display)
+ collection.update_scalarmappable()
+ faces = collection.get_facecolors()
+ edges = collection.get_edgecolors()
+ if isinstance(edges, str) or len(edges) == 0:
+ edges = np.zeros((1, 4))
+ sizes = collection.get_sizes()
+ widths = np.atleast_1d(collection.get_linewidths())
+ marker = _marker_name(paths[0])
+
+ def rows(values):
+ values = np.asarray(values)
+ if len(values) == 0:
+ return np.zeros((count, 4))
+ return values[np.arange(count) % len(values)]
+
+ faces, edges = rows(faces), rows(edges)
+ # marker sizes are areas in pt^2; marks are sized by their diameter
+ if len(sizes):
+ sizes = np.sqrt(np.asarray(sizes, dtype=float))[
+ np.arange(count) % len(sizes)
+ ]
+ else:
+ sizes = np.full(count, 6.0)
+ widths = widths[np.arange(count) % len(widths)]
+ # one plot per style: colors rounded to 1/63, sizes to quarter points
+ keys = np.column_stack(
+ [
+ np.round(faces * 63),
+ np.round(edges * 63),
+ np.round(sizes * 4),
+ np.round(widths * 4),
+ ]
+ )
+ unique, inverse = np.unique(keys, axis=0, return_inverse=True)
+ inverse = np.ravel(inverse)
+ if len(unique) > self.options["max_items"]:
+ return None
+ label = self._legend_label(collection)
+ groups = []
+ for group in range(len(unique)):
+ members = np.flatnonzero((inverse == group) & inside)
+ if not len(members):
+ continue
+ first = members[0]
+ options = ["only marks"] + self._mark(
+ marker,
+ sizes[first],
+ faces[first],
+ edges[first],
+ widths[first],
+ )
+ groups.append((data[members], options))
+
+ def emit():
+ for index, (points, options) in enumerate(groups):
+ self._add(points, options, label if index == 0 else "")
+
+ return emit
+
+ def _line_collection(self, collection):
+ offsets = collection.get_offsets()
+ if len(collection.get_transforms()) or (len(offsets) and np.any(offsets)):
+ return None
+ collection.update_scalarmappable()
+ segments = collection.get_segments()
+ if not segments:
+ return lambda: None
+ transform = collection.get_transform()
+ colors = collection.get_colors()
+ widths = np.atleast_1d(collection.get_linewidths())
+ styles = collection.get_linestyles()
+ alpha = None # the colors include the collection's alpha
+ # consecutive segments of one style become one plot, separated by nan
+ runs = []
+ for index, segment in enumerate(segments):
+ if len(segment) == 0:
+ continue
+ color = colors[index % len(colors)] if len(colors) else (0, 0, 0, 0)
+ width = widths[index % len(widths)]
+ style = styles[index % len(styles)] if len(styles) else (0, None)
+ key = (tuple(np.round(color, 4)), float(width), str(style))
+ display = self._clip_line(
+ transform.transform(np.asarray(segment, dtype=float))
+ )
+ if not len(display):
+ continue
+ if runs and runs[-1][0] == key:
+ runs[-1][1].append(display)
+ else:
+ runs.append((key, [display], color, width, style))
+ if len(runs) > self.options["max_items"]:
+ return None
+ pieces = []
+ for key, displays, color, width, style in runs:
+ stroke = self._stroke(color, width, _dash(style), alpha)
+ if stroke is None:
+ continue
+ joined = []
+ for display in displays:
+ joined.extend([display, [[np.nan, np.nan]]])
+ pieces.append(
+ (self._to_data(np.concatenate(joined[:-1])), stroke + ["mark=none"])
+ )
+ if sum(len(data) for data, _ in pieces) > self.options["max_points"]:
+ return None
+ label = self._legend_label(collection)
+
+ def emit():
+ for index, (data, options) in enumerate(pieces):
+ self._add(data, options, label if index == 0 else "")
+
+ return emit
+
+ def _polygons(self, collection):
+ offsets = collection.get_offsets()
+ if len(collection.get_transforms()) or (len(offsets) and np.any(offsets)):
+ return None
+ paths = collection.get_paths()
+ if len(paths) > self.options["max_items"]:
+ return None
+ collection.update_scalarmappable()
+ transform = collection.get_transform()
+ faces = collection.get_facecolors()
+ edges = collection.get_edgecolors()
+ widths = np.atleast_1d(collection.get_linewidths())
+ styles = collection.get_linestyles()
+ pieces = []
+ for index, path in enumerate(paths):
+ face = faces[index % len(faces)] if len(faces) else "none"
+ edge = edges[index % len(edges)] if len(edges) else "none"
+ width = widths[index % len(widths)]
+ style = styles[index % len(styles)] if len(styles) else (0, None)
+ options = self._area_options(face, edge, width, _dash(style))
+ if options is None:
+ continue
+ for polygon in transform.transform_path(path).to_polygons():
+ polygon = self._clip_polygon(polygon)
+ if len(polygon):
+ pieces.append((self._to_data(polygon), options))
+ label = self._legend_label(collection)
+
+ def emit():
+ for index, (data, options) in enumerate(pieces):
+ self._add(data, options, label if index == 0 else "", cycle=True)
+
+ return emit
+
+ def _area_options(self, face, edge, width, dash, alpha=None):
+ fill = self._fill(face, alpha)
+ stroke = self._stroke(edge, width, dash, alpha)
+ if fill is None and stroke is None:
+ return None
+ options = list(fill or ["fill=none"])
+ options.extend(stroke or ["draw=none"])
+ options.extend(["mark=none", "area legend"])
+ return options
+
+ def _patch(self, patch, clip=None):
+ if clip is None:
+ clip = patch.get_clip_on()
+ transform = patch.get_transform()
+ path = patch.get_path()
+ fillable = None
+ if isinstance(patch, FancyArrowPatch):
+ try:
+ paths, fillable = patch._get_path_in_displaycoord()
+ except Exception: # pragma: no cover - private Matplotlib API
+ return None
+ if not np.iterable(fillable):
+ paths, fillable = [paths], [fillable]
+ transform = None
+ else:
+ paths, fillable = [path], [True]
+ # the patch's colors include its alpha
+ face = patch.get_facecolor() if patch.get_fill() else "none"
+ dash = _dash(patch._dash_pattern)
+ pieces = []
+ for sub_path, can_fill in zip(paths, fillable):
+ if transform is not None:
+ sub_path = transform.transform_path(sub_path)
+ polygons = sub_path.to_polygons(closed_only=False)
+ options = self._area_options(
+ face if can_fill else "none",
+ patch.get_edgecolor(),
+ patch.get_linewidth(),
+ dash,
+ )
+ if options is None:
+ continue
+ for polygon in polygons:
+ closed = can_fill and len(polygon) > 2
+ polygon = (
+ self._clip_polygon(polygon) if closed else self._clip_line(polygon)
+ )
+ if len(polygon):
+ pieces.append((self._to_data(polygon), options, closed))
+ label = self._legend_label(patch)
+
+ def emit():
+ for index, (data, options, closed) in enumerate(pieces):
+ self._add(
+ data, options, label if index == 0 else "", cycle=closed, clip=clip
+ )
+
+ return emit
+
+ def _text(self, text):
+ if not text.get_text().strip():
+ return lambda: None
+ if isinstance(text, Annotation):
+ text.update_positions(self.figure._get_renderer())
+ arrow = (
+ getattr(text, "arrow_patch", None) if isinstance(text, Annotation) else None
+ )
+ arrow_emit = None
+ if arrow is not None and arrow.get_visible():
+ arrow_emit = self._patch(arrow, clip=False)
+ if arrow_emit is None:
+ return None
+ display = _text_display_position(text)
+ if not self.figure.bbox.contains(*display):
+ return arrow_emit or (lambda: None)
+ fx, fy = self.ax.transAxes.inverted().transform(display)
+ node = _text_node(text, self.colors, f"(axis description cs:{fx:.5g},{fy:.5g})")
+
+ def emit():
+ if arrow_emit is not None:
+ arrow_emit()
+ self.axis.add_raw(node)
+
+ return emit
+
+
+# artists drawn as images without a warning: images, meshes, filled contours,
+# quivers and the like have no better form in pgfplots
+_ALWAYS_RASTER = (AxesImage, Collection)
+
+
+def _marker_name(path) -> str:
+ for marker in _MARKS:
+ style = MarkerStyle(marker)
+ candidate = style.get_path().transformed(style.get_transform())
+ if candidate.vertices.shape == path.vertices.shape and np.allclose(
+ candidate.vertices, path.vertices
+ ):
+ return marker
+ return "o"
+
+
+def _with_gaps(path):
+ """The vertices of a path as one array, its separate pieces joined by nan."""
+ out = []
+ for vertex, code in zip(path.vertices, path.codes):
+ if code == Path.STOP:
+ break
+ if code == Path.MOVETO and out:
+ out.append((np.nan, np.nan))
+ out.append(tuple(vertex))
+ return np.asarray(out, dtype=float).reshape(-1, 2)
diff --git a/src/maxplotlib/backends/tikzfigure/text.py b/src/maxplotlib/backends/tikzfigure/text.py
new file mode 100644
index 0000000..616fb9c
--- /dev/null
+++ b/src/maxplotlib/backends/tikzfigure/text.py
@@ -0,0 +1,183 @@
+"""Matplotlib text as LaTeX for the tikzfigure backend.
+
+Matplotlib text is plain text with optional ``$...$`` mathtext. Plain text
+is escaped for LaTeX (``_``, ``%``, ``&`` and so on are literal in
+Matplotlib); mathtext is passed on, being nearly LaTeX already, without
+Matplotlib's own commands such as ``\\mathdefault``. Unicode symbols
+(``ω``, ``−``, ``°``) become math, which pdflatex can typeset.
+"""
+
+import re
+
+# Unicode characters pdflatex does not typeset by default, as math.
+_UNICODE_MATH = {
+ "α": r"\alpha",
+ "β": r"\beta",
+ "γ": r"\gamma",
+ "δ": r"\delta",
+ "ε": r"\varepsilon",
+ "ϵ": r"\epsilon",
+ "ζ": r"\zeta",
+ "η": r"\eta",
+ "θ": r"\theta",
+ "ϑ": r"\vartheta",
+ "ι": r"\iota",
+ "κ": r"\kappa",
+ "λ": r"\lambda",
+ "μ": r"\mu",
+ "µ": r"\mu",
+ "ν": r"\nu",
+ "ξ": r"\xi",
+ "π": r"\pi",
+ "ρ": r"\rho",
+ "σ": r"\sigma",
+ "τ": r"\tau",
+ "υ": r"\upsilon",
+ "φ": r"\varphi",
+ "ϕ": r"\phi",
+ "χ": r"\chi",
+ "ψ": r"\psi",
+ "ω": r"\omega",
+ "Γ": r"\Gamma",
+ "Δ": r"\Delta",
+ "Θ": r"\Theta",
+ "Λ": r"\Lambda",
+ "Ξ": r"\Xi",
+ "Π": r"\Pi",
+ "Σ": r"\Sigma",
+ "Φ": r"\Phi",
+ "Ψ": r"\Psi",
+ "Ω": r"\Omega",
+ "−": "-",
+ "±": r"\pm",
+ "∓": r"\mp",
+ "×": r"\times",
+ "·": r"\cdot",
+ "÷": r"\div",
+ "°": r"^\circ",
+ "≈": r"\approx",
+ "∼": r"\sim",
+ "≤": r"\leq",
+ "≥": r"\geq",
+ "≠": r"\neq",
+ "∝": r"\propto",
+ "∞": r"\infty",
+ "∂": r"\partial",
+ "∇": r"\nabla",
+ "∫": r"\int",
+ "∑": r"\sum",
+ "√": r"\surd",
+ "‖": r"\|",
+ "∥": r"\parallel",
+ "⊥": r"\perp",
+ "→": r"\rightarrow",
+ "←": r"\leftarrow",
+ "↔": r"\leftrightarrow",
+ "⟨": r"\langle",
+ "⟩": r"\rangle",
+ "ℏ": r"\hbar",
+ "ℓ": r"\ell",
+ "′": r"\prime",
+ "⊙": r"\odot",
+ "⊗": r"\otimes",
+ "⁰": "^{0}",
+ "¹": "^{1}",
+ "²": "^{2}",
+ "³": "^{3}",
+ "⁴": "^{4}",
+ "⁵": "^{5}",
+ "⁶": "^{6}",
+ "⁷": "^{7}",
+ "⁸": "^{8}",
+ "⁹": "^{9}",
+ "⁻": "^{-}",
+ "₀": "_{0}",
+ "₁": "_{1}",
+ "₂": "_{2}",
+ "₃": "_{3}",
+}
+
+_TEXT_ESCAPES = {
+ "\\": r"\textbackslash{}",
+ "&": r"\&",
+ "%": r"\%",
+ "#": r"\#",
+ "_": r"\_",
+ "{": r"\{",
+ "}": r"\}",
+ "~": r"\textasciitilde{}",
+ "^": r"\textasciicircum{}",
+ "$": r"\$",
+}
+
+# mathtext commands that LaTeX lacks, and what they mean there
+_MATH_COMMANDS = [
+ (re.compile(r"\\mathdefault\{([^{}]*)\}"), r"\1"),
+ (re.compile(r"\\mathdefault\b"), ""),
+ (re.compile(r"\\degree\b"), r"{^\\circ}"),
+ (re.compile(r"\\AA\b"), r"\\text{\\AA}"),
+]
+
+
+def _math(source: str) -> str:
+ for pattern, replacement in _MATH_COMMANDS:
+ source = pattern.sub(replacement, source)
+ out = []
+ for char in source:
+ symbol = _UNICODE_MATH.get(char)
+ if symbol is None:
+ out.append(char)
+ elif symbol.startswith("\\") and symbol[-1:].isalpha():
+ out.append(symbol + " ")
+ else:
+ out.append(symbol)
+ return "".join(out)
+
+
+def _text(source: str) -> str:
+ out = []
+ for char in source:
+ if char in _TEXT_ESCAPES:
+ out.append(_TEXT_ESCAPES[char])
+ elif char in _UNICODE_MATH:
+ out.append(f"${_UNICODE_MATH[char]}$")
+ elif char == "\n":
+ out.append(r"\\")
+ else:
+ out.append(char)
+ return "".join(out)
+
+
+def latex(text) -> str:
+ r"""Matplotlib text as LaTeX: plain parts escaped, ``$...$`` parts as math.
+
+ An odd number of unescaped ``$`` makes Matplotlib show the text as it
+ is, so then every part is plain text.
+
+ Examples
+ --------
+ >>> latex(r"growth rate $\gamma/\omega_{ci}$ (50% fit)")
+ 'growth rate $\\gamma/\\omega_{ci}$ (50\\% fit)'
+ >>> latex(r"$\mathdefault{0.5}$")
+ '$0.5$'
+ >>> latex("ω = 2.5")
+ '$\\omega$ = 2.5'
+ """
+ if text is None:
+ return ""
+ text = str(text)
+ parts = re.split(r"(? bool:
+ """Whether ``text`` has more than one line, needing ``align`` in TikZ."""
+ return "\n" in str(text or "")
diff --git a/src/maxplotlib/canvas/canvas.py b/src/maxplotlib/canvas/canvas.py
index 2a5fcca..fb08d64 100644
--- a/src/maxplotlib/canvas/canvas.py
+++ b/src/maxplotlib/canvas/canvas.py
@@ -18,13 +18,8 @@
from maxplotlib.backends.plotext import PlotextFigure, create_plotext_figure
from maxplotlib.colors.colors import Color
from maxplotlib.linestyle.linestyle import Linestyle
-from maxplotlib.subfigure.line_plot import (
- _TIKZ_SUPPORTED_PLOT_TYPES,
- LinePlot,
- _tikz_error_bounds,
- _tikz_step_coordinates,
- _tikz_style_kwargs,
-)
+from maxplotlib.subfigure.line_plot import LinePlot
+from maxplotlib.utils import xarray_support
from maxplotlib.utils.options import Backends
@@ -370,6 +365,10 @@ def __init__(
self._supylabel_kwargs: dict = {}
self._subplots_adjust_kwargs: dict = {}
self._tight_layout_kwargs: dict | None = None
+ self._hide_empty_subplots = False
+ # Set by Canvas.facet: merge Plotly legend entries across subplots and
+ # leave room for the figure title above the subplot titles.
+ self._facet = False
self._set_tight_layout = None
self._align_labels = False
self._align_titles = False
@@ -393,6 +392,40 @@ def __init__(
# Factory
# ------------------------------------------------------------------
+ @classmethod
+ def from_matplotlib(
+ cls, source, *, strict=False, trusted=False, fallback="native", **canvas_kwargs
+ ):
+ """Snapshot a Matplotlib figure, axes, or rectangular axes array.
+
+ Portable geometry becomes independent plot entries. With the default
+ ``fallback="native"``, other built-in artists retain detached native
+ geometry for Matplotlib, including transforms and decorations. These
+ entries reject rendering on backends that cannot represent them.
+ ``fallback="skip"`` instead warns about unsupported artists;
+ ``fallback="raster"`` explicitly flattens the selection into an image.
+
+ ``strict=True`` raises on import losses; it does not promise identical
+ output on every backend. ``canvas.import_report`` contains artist
+ identities, severities, representations and backend restrictions.
+
+ Paths, pickle bytes and binary streams require ``trusted=True`` before
+ reading: unpickling can execute arbitrary code. Serialized figures are
+ only suitable for trusted producers and matching Matplotlib versions.
+ ``canvas_kwargs`` override figure defaults. Two-dimensional input
+ arrays define their own slot order; other inputs preserve source layout.
+ """
+ from maxplotlib.backends.matplotlib.importer import import_matplotlib
+
+ return import_matplotlib(
+ cls,
+ source,
+ strict=strict,
+ trusted=trusted,
+ fallback=fallback,
+ **canvas_kwargs,
+ )
+
@classmethod
def subplots(
cls,
@@ -452,6 +485,171 @@ def subplots(
return canvas, [row[0] for row in axes]
return canvas, axes
+ _FACET_KINDS = {
+ "plot": "plot",
+ "scatter": "scatter",
+ "pcolormesh": "pcolormesh",
+ "imshow": "add_imshow",
+ "contour": "contour",
+ "contourf": "contourf",
+ }
+
+ @classmethod
+ def facet(
+ cls,
+ da,
+ col=None,
+ row=None,
+ col_wrap: int | None = None,
+ kind: str = "pcolormesh",
+ sharey: bool = True,
+ canvas_kwargs: dict | None = None,
+ **kwargs,
+ ):
+ """
+ Create a Canvas with one subplot per value of a DataArray dimension.
+
+ Parameters:
+ da (xarray.DataArray): The data. After removing ``col``/``row``, it
+ must have the dimensions ``kind`` needs: 1-D for ``"plot"`` and
+ ``"scatter"`` (2-D with ``hue=``), 2-D for the others.
+ col, row (str): Dimensions to lay out across columns and rows.
+ col_wrap (int): Wrap a ``col`` facet after this many columns.
+ kind (str): ``"plot"``, ``"scatter"``, ``"pcolormesh"``, ``"imshow"``,
+ ``"contour"`` or ``"contourf"``.
+ sharey (bool): For ``"plot"``/``"scatter"``, give every subplot the
+ value range of the whole array (on the x-axis with ``ycoord=``).
+ With ``False`` each subplot scales its own axis and keeps its y
+ tick labels.
+ canvas_kwargs (dict): Forwarded to the Canvas constructor.
+ **kwargs: Forwarded to each subplot's ``kind`` method.
+
+ Color-mapped kinds share one color scale (``vmin``/``vmax`` default to
+ the data range, and contour levels are shared) and one colorbar for
+ the whole figure, which ``add_colorbar=False`` turns off. Axis and
+ tick labels are kept on the outer subplots only. Each subplot is
+ titled with the coordinate values that differ between subplots, and
+ the figure with the ones they share.
+
+ Returns:
+ (canvas, axes): The Canvas and a 2-D list of LinePlots, with ``None``
+ where a wrapped grid has no subplot.
+
+ Examples:
+ >>> canvas, axes = Canvas.facet(da, col="t", col_wrap=3)
+ >>> canvas, axes = Canvas.facet(da, row="species", kind="plot")
+ """
+ if not xarray_support.is_dataarray(da):
+ raise TypeError("facet() needs an xarray.DataArray")
+ if kind not in cls._FACET_KINDS:
+ raise ValueError(
+ f"kind must be one of {sorted(cls._FACET_KINDS)}, got {kind!r}"
+ )
+ if col is None and row is None:
+ raise ValueError("facet() needs col= and/or row=")
+ if col == row:
+ raise ValueError("col and row must be different dimensions")
+ for dim in (col, row):
+ if dim is not None and dim not in da.dims:
+ raise ValueError(f"{dim!r} is not a dimension of {da.dims}")
+ if col_wrap is not None and (col is None or row is not None):
+ raise ValueError("col_wrap= needs col= and no row=")
+
+ ncol_values = da.sizes[col] if col is not None else 1
+ nrow_values = da.sizes[row] if row is not None else 1
+ if col_wrap is not None:
+ ncols = max(1, min(col_wrap, ncol_values))
+ nrows = -(-ncol_values // ncols)
+ panels = [(j // ncols, j % ncols, {col: j}) for j in range(ncol_values)]
+ else:
+ ncols, nrows = ncol_values, nrow_values
+ panels = [
+ (i, j, {d: k for d, k in ((row, i), (col, j)) if d is not None})
+ for i in range(nrows)
+ for j in range(ncols)
+ ]
+
+ mapped = kind not in ("plot", "scatter")
+ sharey = sharey or mapped
+ add_colorbar = kwargs.pop("add_colorbar", kind != "contour")
+ values = xarray_support.magnitude(da)
+ if mapped:
+ # Shared limits for every panel, with xarray's color defaults.
+ xarray_support.color_limits(values, kwargs)
+ kwargs.setdefault("vmin", float(np.nanmin(values)))
+ kwargs.setdefault("vmax", float(np.nanmax(values)))
+ if kind in ("contour", "contourf") and "levels" not in kwargs:
+ from matplotlib.ticker import MaxNLocator
+
+ kwargs["levels"] = MaxNLocator(10).tick_values(
+ kwargs["vmin"], kwargs["vmax"]
+ )
+ # One colorbar for the figure, added below; "colorbar" hides the
+ # per-trace scale Plotly would otherwise show.
+ kwargs["add_colorbar"] = False
+ kwargs["colorbar"] = False
+
+ canvas_kwargs = dict(canvas_kwargs or {})
+ if not {"subplot_spacing", "gridspec_kw"} & canvas_kwargs.keys():
+ canvas_kwargs["subplot_spacing"] = SubplotSpacing(
+ wspace=0.1 if sharey else 0.3, hspace=0.3
+ )
+ canvas = cls(nrows=nrows, ncols=ncols, **canvas_kwargs)
+ canvas._hide_empty_subplots = True
+ canvas._facet = True
+ # Coordinates shared by every panel go in the figure title, so the
+ # panel titles only show what differs between them.
+ common = [name for name, coord in da.coords.items() if coord.ndim == 0]
+ common_title = xarray_support.title(da)
+ if common_title:
+ canvas.suptitle(common_title)
+ axes = [[None] * ncols for _ in range(nrows)]
+ method = cls._FACET_KINDS[kind]
+ for r, c, selection in panels:
+ subplot = canvas.add_subplot(row=r, col=c)
+ panel = da.isel(selection)
+ subplot.set_title(
+ xarray_support.title(panel, exclude=common)
+ or ", ".join(f"{dim} = {index}" for dim, index in selection.items())
+ )
+ getattr(subplot, method)(panel, **kwargs)
+ axes[r][c] = subplot
+
+ vertical = "ycoord" in kwargs
+ if not mapped and sharey:
+ low, high = float(np.nanmin(values)), float(np.nanmax(values))
+ margin = 0.05 * (high - low)
+ for r, c, _ in panels:
+ subplot = axes[r][c]
+ tick_params = {}
+ if r + 1 < nrows and axes[r + 1][c] is not None:
+ subplot._xlabel = None
+ tick_params["labelbottom"] = False
+ if c > 0 and sharey:
+ subplot._ylabel = None
+ tick_params["labelleft"] = False
+ if tick_params:
+ subplot.tick_params(**tick_params)
+ if not mapped and sharey and vertical:
+ subplot.set_xlim(low - margin, high + margin)
+ elif not mapped and sharey:
+ subplot.set_ylim(low - margin, high + margin)
+ if (r, c) != panels[0][:2]:
+ subplot._legend = False
+ if mapped and add_colorbar:
+ first = axes[panels[0][0]][panels[0][1]]
+ first._add(
+ {
+ "label": xarray_support.value_label(da),
+ "layer": 0,
+ "plot_type": "colorbar",
+ "kwargs": {},
+ "span_figure": True,
+ },
+ 0,
+ )
+ return canvas, axes
+
@property
def _subplot_dict(self):
return self._subplots
@@ -575,7 +773,7 @@ def _get_or_create_subplot(self, row, col):
def scatter(
self,
x,
- y,
+ y=None,
layer=0,
row: int | None = None,
col: int | None = None,
@@ -590,6 +788,8 @@ def scatter(
layer (int): Layer index (default 0).
row, col (int): Subplot position (default top-left).
**kwargs: Forwarded to the backend (e.g., color, marker, s, label).
+
+ ``scatter(da)`` accepts an ``xarray.DataArray`` like ``plot(da)``.
"""
sp = self._get_or_create_subplot(row, col)
sp.scatter(x, y, layer=layer, **kwargs)
@@ -796,40 +996,44 @@ def eventplot(
def contour(
self,
x,
- y,
- z,
+ y=None,
+ z=None,
layer=0,
row: int | None = None,
col: int | None = None,
**kwargs,
):
- """Add contour lines to a subplot."""
+ """Add contour lines to a subplot; ``contour(da)`` takes a DataArray."""
self._get_or_create_subplot(row, col).contour(x, y, z, layer=layer, **kwargs)
def contourf(
self,
x,
- y,
- z,
+ y=None,
+ z=None,
layer=0,
row: int | None = None,
col: int | None = None,
**kwargs,
):
- """Add filled contours to a subplot."""
+ """Add filled contours to a subplot; ``contourf(da)`` takes a DataArray."""
self._get_or_create_subplot(row, col).contourf(x, y, z, layer=layer, **kwargs)
def pcolormesh(
self,
x,
- y,
- z,
+ y=None,
+ z=None,
layer=0,
row: int | None = None,
col: int | None = None,
**kwargs,
):
- """Add a pseudocolor mesh to a subplot."""
+ """Add a pseudocolor mesh to a subplot.
+
+ ``pcolormesh(da)`` accepts a 2-D ``xarray.DataArray``; see
+ :meth:`LinePlot.pcolormesh`.
+ """
self._get_or_create_subplot(row, col).pcolormesh(x, y, z, layer=layer, **kwargs)
def hexbin(
@@ -1559,7 +1763,7 @@ def imshow(
col: int | None = None,
**kwargs,
):
- """Add an image/matrix plot to a subplot."""
+ """Add an image/matrix plot to a subplot; ``imshow(da)`` takes a DataArray."""
self._get_or_create_subplot(row, col).add_imshow(data, layer=layer, **kwargs)
def add_image(
@@ -1867,7 +2071,17 @@ def savefig(
layers: list | None = None,
layer_by_layer: bool = False,
verbose: bool = False,
+ include_plotlyjs: bool | str = True,
):
+ """Render and save the canvas.
+
+ ``include_plotlyjs`` only applies to the Plotly backend when saving
+ to an ``.html``/``.htm`` file (see ``fig.write_html`` in the Plotly
+ docs). It defaults to ``True``, which bundles plotly.js into the
+ file so it works offline. Pass ``"cdn"`` to instead reference
+ plotly.js from a CDN, producing a much smaller file that requires
+ network access to render.
+ """
filename_no_extension, extension = os.path.splitext(filename)
if backend == "matplotlib":
if layer_by_layer:
@@ -1942,7 +2156,9 @@ def savefig(
savefig=False,
layers=layers,
)
- self._save_plotly(fig, full_filepath)
+ self._save_plotly(
+ fig, full_filepath, include_plotlyjs=include_plotlyjs
+ )
if verbose:
print(f"Saved {full_filepath}")
else:
@@ -1956,7 +2172,7 @@ def savefig(
savefig=False,
layers=layers,
)
- self._save_plotly(fig, full_filepath)
+ self._save_plotly(fig, full_filepath, include_plotlyjs=include_plotlyjs)
if verbose:
print(f"Saved {full_filepath}")
elif backend == "tikzfigure":
@@ -2028,7 +2244,7 @@ def _render(
verbose=verbose,
)
elif backend == "tikzfigure":
- return self.plot_tikzfigure(savefig=savefig, verbose=verbose)
+ return self.plot_tikzfigure(savefig=savefig, layers=layers, verbose=verbose)
else:
raise ValueError(f"Invalid backend: {backend}")
@@ -2036,10 +2252,19 @@ def plot(self, *args, backend=None, **kwargs):
"""Add a line, or render when called with backend options.
``canvas.plot(x, y, **style)`` is the convenient direct plotting form.
+ ``canvas.plot(da)`` plots a 1-D ``xarray.DataArray`` against its
+ coordinate, labelling the axes from its attributes; ``hue=``
+ draws a 2-D one as one line per value of ``dim``.
Rendering is named explicitly by ``canvas.render(...)``; the legacy
``canvas.plot(backend=...)`` form remains supported.
"""
explicit_render = backend is not None or (args and isinstance(args[0], str))
+ if len(args) == 1 and xarray_support.is_dataarray(args[0]):
+ layer = kwargs.pop("layer", 0)
+ row = kwargs.pop("row", None)
+ col = kwargs.pop("col", None)
+ self._get_or_create_subplot(row, col).plot(args[0], layer=layer, **kwargs)
+ return self
if args and not isinstance(args[0], str):
if len(args) < 2:
raise TypeError("plot(x, y) requires both x and y data")
@@ -2290,14 +2515,62 @@ def plot_matplotlib(
if verbose:
print(f"Created Matplotlib figure and axes with shape {axes.shape}")
+ if hasattr(self, "_import_layout"):
+ for row in range(self.nrows):
+ for col in range(self.ncols):
+ if (row, col) not in self._subplot_dict:
+ axes[row, col].remove()
+ axes[row, col] = None
+ for slot, bounds in self._import_layout.items():
+ axes[slot].set_position(bounds)
+ if hasattr(self, "_import_layout_specs"):
+ for slot, spec in self._import_layout_specs.clone(fig=fig).items():
+ if spec is not None:
+ axes[slot].set_subplotspec(spec)
+ else:
+ bounds = axes[slot].get_position(original=True).bounds
+ axes[slot].remove()
+ axes[slot] = fig.add_axes(bounds)
+ if self._import_layout_engine is not None:
+ import copy
+
+ fig.set_layout_engine(copy.deepcopy(self._import_layout_engine))
+ for slot, in_layout in self._import_in_layout.items():
+ axes[slot].set_in_layout(in_layout)
+ for direction, parent, child in self._import_shared_axes:
+ getattr(axes[child], "share" + direction)(axes[parent])
+ if hasattr(self, "_import_figure_style"):
+ fig.set(**self._import_figure_style)
+ fig.patch = self._import_figure_patch.clone(fig=fig)
+ transform = fig.patch.get_transform()
+ fig._set_artist_props(fig.patch)
+ fig.patch.set_transform(transform)
for (row, col), subplot in self._subplot_dict.items():
ax = axes[row][col]
+ if hasattr(subplot, "_import_projection"):
+ position = ax.get_position().bounds
+ ax.remove()
+ ax = subplot._import_projection.clone(fig=fig)
+ fig.add_axes(ax)
+ ax.set_position(position)
+ axes[row, col] = ax
if isinstance(subplot, TikzFigure):
plot_matplotlib(subplot, ax, layers=layers)
+ ax.grid(False)
else:
subplot.plot_matplotlib(ax, layers=layers)
- # ax.set_title(f"Subplot ({row}, {col})")
- ax.grid()
+
+ if self._hide_empty_subplots:
+ for row in range(self.nrows):
+ for col in range(self.ncols):
+ if (row, col) not in self._subplot_dict:
+ axes[row][col].set_visible(False)
+ figure_axes = [ax for ax in fig.axes if ax.get_visible()]
+ for subplot in self._subplot_dict.values():
+ figure_colorbar = getattr(subplot, "_figure_colorbar", None)
+ if figure_colorbar is not None and figure_colorbar[0] is not None:
+ mappable, label = figure_colorbar
+ fig.colorbar(mappable, ax=figure_axes, label=label)
if verbose:
print("Finished plotting subplots.")
@@ -2311,6 +2584,10 @@ def plot_matplotlib(
fig.supxlabel(self._supxlabel, **self._supxlabel_kwargs)
if self._supylabel:
fig.supylabel(self._supylabel, **self._supylabel_kwargs)
+ if self._facet and self._suptitle and "top" not in self._subplots_adjust_kwargs:
+ # About four font heights: the figure title plus subplot titles.
+ room = 4 * self.fontsize / 72
+ fig.subplots_adjust(top=max(0.5, 1 - room / fig.get_figheight()))
if self._subplots_adjust_kwargs:
fig.subplots_adjust(**self._subplots_adjust_kwargs)
if self._tight_layout_kwargs is not None:
@@ -2343,6 +2620,43 @@ def plot_matplotlib(
twin_axis = axes[row][col].twiny()
twin_subplot.plot_matplotlib(twin_axis, layers=layers)
self._matplotlib_twiny_axes[(row, col)] = twin_axis
+ if hasattr(self, "_import_layout"):
+ from maxplotlib.backends.matplotlib.import_state import apply_axis_state
+
+ for slot, direction, subplot in self._import_extra_twins:
+ twin_axis = getattr(axes[slot], "twin" + direction)()
+ subplot.plot_matplotlib(twin_axis, layers=layers)
+ for specification in self._import_colorbars:
+ if specification["standalone"] is not None:
+ import copy
+
+ from matplotlib.cm import ScalarMappable
+
+ mappable = ScalarMappable(
+ **copy.deepcopy(specification["standalone"])
+ )
+ else:
+ artists = specification["target"]._import_rendered_artists.get(
+ specification["artist_id"], []
+ )
+ mappable = next(
+ (artist for artist in artists if hasattr(artist, "get_cmap")),
+ None,
+ )
+ if mappable is None:
+ continue # Its layer was excluded from this render.
+ cax = fig.add_axes(specification["position"])
+ colorbar = fig.colorbar(mappable, cax=cax, **specification["kwargs"])
+ colorbar.set_label(
+ specification["label"], **specification["label_style"]
+ )
+ colorbar.set_ticks(specification["ticks"])
+ colorbar.formatter = specification["formatter"].clone(cax)
+ colorbar.update_ticks()
+ apply_axis_state(cax, specification["state"])
+ for snapshot in self._import_figure_artists:
+ artist = snapshot.clone(fig=fig)
+ fig.add_artist(artist)
if matplotlib_customizations is not None:
_apply_matplotlib_customizations(fig, axes, matplotlib_customizations)
if matplotlib_postprocess is not None:
@@ -2351,335 +2665,123 @@ def plot_matplotlib(
matplotlib_postprocess(fig, axes)
return fig, axes
+ def _validate_import_backend(self, backend, *, allow_unsupported=False):
+ """Never silently drop native imported axes or figure decorations."""
+ if not hasattr(self, "import_report") or allow_unsupported:
+ return
+ reasons = []
+ if getattr(self, "_import_extra_twins", []):
+ reasons.append("multiple twin axes")
+ if getattr(self, "_import_figure_artists", []):
+ reasons.append("native figure decorations")
+ if getattr(self, "_import_colorbars", []):
+ reasons.append("native colorbars")
+ subplots = (
+ list(self._subplot_dict.values())
+ + list(self._twinx_subplots.values())
+ + list(self._twiny_subplots.values())
+ )
+ for subplot in subplots:
+ if hasattr(subplot, "_import_projection"):
+ reasons.append("native projections")
+ if getattr(subplot, "_import_child_axes", []):
+ reasons.append("inset/secondary axes")
+ for state in (
+ getattr(subplot, "_import_axis_state", {}).get("axes", {}).values()
+ ):
+ scale = state["scale"].payload
+ if scale.name not in ("linear", "log"):
+ reasons.append("native axis scales")
+ if reasons:
+ raise NotImplementedError(
+ f"{backend} cannot render these imported features: {', '.join(sorted(set(reasons)))}. "
+ "Use Matplotlib or import with fallback='raster'."
+ )
+
def plot_tikzfigure(
self,
savefig: bool = False,
verbose: bool = False,
+ *,
+ layers: list | None = None,
+ raster_dpi: float = 300,
+ max_markers: int = 2000,
+ max_items: int = 500,
+ max_points: int = 20000,
+ precision: int = 6,
) -> TikzFigure:
- """
- Generate a TikZ figure from subplots.
+ """Render the canvas as a TikZ/pgfplots figure.
- For now, returns the first subplot's TikzFigure.
- Full multi-subplot support requires TikzFigure's subfigure_axis API.
+ The canvas is drawn with Matplotlib, off screen, and the drawn figure
+ is converted with
+ :func:`~maxplotlib.backends.tikzfigure.figure_to_tikz`: every subplot
+ becomes a pgfplots axis at the same place and size, with its labels,
+ ticks, legend and colorbar, lines, markers, bars, fills and text as
+ pgfplots code, and meshes, images and other artists without a vector
+ counterpart as images Matplotlib renders (``\\addplot graphics``).
+ Every layout, twin axes and imported figure therefore converts.
Parameters:
- verbose (bool): If True, print debug information.
+ savefig (bool): Unused; kept for the other backends' signature.
+ verbose (bool): If True, print progress.
+ layers (list): Draw only these layers, as with Matplotlib.
+ raster_dpi (float): Resolution of the parts drawn as images.
+ max_markers (int): Scatter plots with more points are drawn as an image.
+ max_items (int): Collections with more differently styled items are
+ drawn as an image.
+ max_points (int): Lines with more points (after simplification) are
+ drawn as an image.
+ precision (int): Significant digits of the written coordinates.
Returns:
- TikzFigure: Figure object that can be shown, saved, or compiled.
+ TikzFigure: Figure object that can be shown, saved (``.tikz`` and
+ ``.tex`` with the images next to them, ``.pdf``, ``.png``) or
+ compiled.
"""
- if verbose:
- print(f"Plotting tikzfigure with {len(self._subplot_dict)} subplot(s)")
-
- if self._twinx_subplots:
- raise NotImplementedError(
- "twinx plots are currently supported only by the matplotlib and plotly backends"
- )
-
- # Check for unsupported layouts
- if self.nrows > 1:
- raise NotImplementedError(
- "Vertical/grid layouts (nrows > 1) are not yet supported for tikzfigure backend. "
- "Use horizontal layouts (1×n) only."
- )
-
- # Validate that at least one subplot exists
- if len(self._subplot_dict) == 0:
- raise ValueError(
- "No subplots to plot. Call add_subplot() or Canvas.subplots() first."
- )
-
- axis_width, axis_height = self._get_tikzfigure_axis_dimensions()
- fig = TikzFigure()
-
- # Add each subplot as a subfigure axis
- for (row, col), line_plot in self._subplot_dict.items():
- if verbose:
- print(f"Plotting subplot at row {row}, col {col}")
-
- # Create subfigure axis with subplot metadata
- ax = fig.subfigure_axis(
- xlabel=line_plot._xlabel or "",
- ylabel=line_plot._ylabel or "",
- xlim=(
- (line_plot._xmin, line_plot._xmax)
- if line_plot._xmin is not None
- else None
- ),
- ylim=(
- (line_plot._ymin, line_plot._ymax)
- if line_plot._ymin is not None
- else None
- ),
- grid=line_plot._grid,
- title=line_plot._title or f"Subplot {col + 1}",
- width=0.45,
- axis_width=axis_width,
- height=axis_height,
- )
+ from maxplotlib.backends.tikzfigure import figure_to_tikz
- # Add each plot line to the subfigure
- for line_data in line_plot.line_data:
- plot_type = line_data.get("plot_type")
- if plot_type not in _TIKZ_SUPPORTED_PLOT_TYPES:
- raise NotImplementedError(
- f"{plot_type} is not supported by the tikzfigure backend"
- )
- if plot_type == "plot":
- # Extract and transform x, y data
- x = (line_data["x"] + line_plot._xshift) * line_plot._xscale
- y = (line_data["y"] + line_plot._yshift) * line_plot._yscale
- kwargs = line_data.get("kwargs", {})
- if verbose:
- print(f"Line {kwargs = }")
- # Add plot to subfigure axis
- ax.add_plot(
- x=x,
- y=y,
- **_tikz_style_kwargs(kwargs),
- )
- elif plot_type == "scatter":
- x = (line_data["x"] + line_plot._xshift) * line_plot._xscale
- y = (line_data["y"] + line_plot._yshift) * line_plot._yscale
- kwargs = _tikz_style_kwargs(line_data.get("kwargs", {}))
- kwargs.setdefault("mark", "*")
- kwargs["line_width"] = 0
- ax.add_plot(x=x, y=y, **kwargs)
- elif plot_type in {"bar", "barh"}:
- source_kwargs = line_data.get("kwargs", {})
- kwargs = _tikz_style_kwargs(source_kwargs)
- kwargs["fill"] = source_kwargs.get("color", "blue")
- kwargs["fill_opacity"] = source_kwargs.get("alpha", 1.0)
- kwargs["line_width"] = source_kwargs.get("linewidth", 0)
- if plot_type == "bar":
- width = source_kwargs.get("width", 0.8)
- for x, height in zip(line_data["x"], line_data["height"]):
- ax.add_plot(
- x=[
- x - width / 2,
- x + width / 2,
- x + width / 2,
- x - width / 2,
- ],
- y=[0, 0, height, height],
- cycle=True,
- **kwargs,
- )
- else:
- height = source_kwargs.get("height", 0.8)
- for y, width in zip(line_data["y"], line_data["width"]):
- ax.add_plot(
- x=[0, width, width, 0],
- y=[
- y - height / 2,
- y - height / 2,
- y + height / 2,
- y + height / 2,
- ],
- cycle=True,
- **kwargs,
- )
- elif plot_type == "fill_between":
- x = line_data["x"]
- y1 = np.asarray(line_data["y1"])
- y2 = np.broadcast_to(line_data["y2"], y1.shape)
- source_kwargs = line_data.get("kwargs", {})
- kwargs = _tikz_style_kwargs(source_kwargs)
- kwargs["fill"] = source_kwargs.get("color", "blue")
- kwargs["fill_opacity"] = source_kwargs.get("alpha", 0.25)
- ax.add_plot(
- x=list(x) + list(x[::-1]),
- y=list(y1) + list(y2[::-1]),
- cycle=True,
- **kwargs,
- )
- elif plot_type == "errorbar":
- x = line_data["x"]
- y = line_data["y"]
- kwargs = _tikz_style_kwargs(line_data.get("kwargs", {}))
- ax.add_plot(x=x, y=y, **kwargs)
- y_bounds = _tikz_error_bounds(line_data["yerr"], y)
- if y_bounds is not None:
- lower, upper = y_bounds
- for xi, low, high in zip(x, y - lower, y + upper):
- ax.add_plot(x=[xi, xi], y=[low, high], **kwargs)
- x_bounds = _tikz_error_bounds(line_data["xerr"], x)
- if x_bounds is not None:
- lower, upper = x_bounds
- for yi, low, high in zip(y, x - lower, x + upper):
- ax.add_plot(x=[low, high], y=[yi, yi], **kwargs)
- elif plot_type in {"step", "stairs"}:
- source_kwargs = line_data.get("kwargs", {})
- if plot_type == "step":
- x = line_data["x"]
- y = line_data["y"]
- where = source_kwargs.get("where", "pre")
- else:
- values = line_data["values"]
- edges = line_data["edges"]
- if edges is None:
- edges = np.arange(len(values) + 1)
- x = edges
- y = np.r_[values, values[-1]]
- where = "post"
- x, y = _tikz_step_coordinates(x, y, where=where)
- ax.add_plot(
- x=x,
- y=y,
- **_tikz_style_kwargs(source_kwargs),
- )
- elif plot_type == "stem":
- x = line_data["x"]
- y = line_data["y"]
- source_kwargs = line_data.get("kwargs", {})
- style = _tikz_style_kwargs(source_kwargs)
- marker_style = dict(style)
- marker_style.update(
- mark=source_kwargs.get("marker", "*"), line_width=0
- )
- ax.add_plot(x=x, y=y, **marker_style)
- for xi, yi in zip(x, y):
- ax.add_plot(x=[xi, xi], y=[0, yi], **style)
- elif plot_type in {"hlines", "vlines"}:
- style = _tikz_style_kwargs(line_data.get("kwargs", {}))
- if plot_type == "hlines":
- for yi, left, right in zip(
- np.atleast_1d(line_data["y"]),
- np.atleast_1d(line_data["xmin"]),
- np.atleast_1d(line_data["xmax"]),
- ):
- ax.add_plot(x=[left, right], y=[yi, yi], **style)
- else:
- for xi, bottom, top in zip(
- np.atleast_1d(line_data["x"]),
- np.atleast_1d(line_data["ymin"]),
- np.atleast_1d(line_data["ymax"]),
- ):
- ax.add_plot(x=[xi, xi], y=[bottom, top], **style)
- elif plot_type in {"axvspan", "axhspan"}:
- source_kwargs = line_data.get("kwargs", {})
- style = _tikz_style_kwargs(source_kwargs)
- style["fill"] = source_kwargs.get("color", "blue")
- style["fill_opacity"] = source_kwargs.get("alpha", 0.2)
- if plot_type == "axvspan":
- xmin, xmax = line_data["xmin"], line_data["xmax"]
- ymin, ymax = line_plot._ymin or 0, line_plot._ymax or 1
- x = [xmin, xmax, xmax, xmin]
- y = [ymin, ymin, ymax, ymax]
- else:
- ymin, ymax = line_data["ymin"], line_data["ymax"]
- xmin, xmax = line_plot._xmin or 0, line_plot._xmax or 1
- x = [xmin, xmax, xmax, xmin]
- y = [ymin, ymin, ymax, ymax]
- ax.add_plot(x=x, y=y, cycle=True, **style)
- elif plot_type == "fill":
- if len(line_data["args"]) < 2:
- raise ValueError("tikzfigure fill requires x and y coordinates")
- x, y = line_data["args"][:2]
- source_kwargs = line_data.get("kwargs", {})
- style = _tikz_style_kwargs(source_kwargs)
- style["fill"] = source_kwargs.get("color", "blue")
- style["fill_opacity"] = source_kwargs.get("alpha", 0.25)
- ax.add_plot(x=x, y=y, cycle=True, **style)
- elif plot_type == "flame_chart":
- labels = line_data["labels"]
- parents = line_data["parents"]
- values = line_data["values"] * line_plot._xscale
- start_times = line_data["start_times"]
- depths = np.zeros(len(labels), dtype=int)
- if start_times is None:
- start_times = np.zeros(len(labels))
- else:
- start_times = (
- start_times + line_plot._xshift
- ) * line_plot._xscale
- for index, parent in enumerate(parents):
- if parent is not None:
- parent_index = (
- parent
- if isinstance(parent, int)
- else labels.index(parent)
- )
- depths[index] = depths[parent_index] + 1
- colors = ["red", "blue", "green", "orange", "purple", "cyan"]
- for index, (start, value) in enumerate(zip(start_times, values)):
- y = depths[index]
- ax.add_plot(
- x=[start, start + value, start + value, start],
- y=[y - 0.4, y - 0.4, y + 0.4, y + 0.4],
- cycle=True,
- fill=colors[y % len(colors)],
- line_width=0,
- )
- elif plot_type == "gantt":
- tasks = line_data["tasks"]
- start_times = (
- line_data["start_times"] + line_plot._xshift
- ) * line_plot._xscale
- durations = line_data["durations"] * line_plot._xscale
- y_positions = np.arange(len(tasks))
- kwargs = line_data.get("kwargs", {})
-
- # Draw horizontal bars for each task as filled rectangles
- for i, (task, start, duration) in enumerate(
- zip(tasks, start_times, durations)
- ):
- x_start = float(start)
- x_end = float(start + duration)
- y_pos = float(y_positions[i])
- bar_height = 0.8
-
- # Create rectangle coordinates for the bar
- x_coords = [x_start, x_end, x_end, x_start, x_start]
- y_coords = [
- y_pos - bar_height / 2,
- y_pos - bar_height / 2,
- y_pos + bar_height / 2,
- y_pos + bar_height / 2,
- y_pos - bar_height / 2,
- ]
-
- # Add as a filled plot
- color = kwargs.get("color", "blue")
- ax.add_plot(
- x=x_coords,
- y=y_coords,
- color=color,
- fill=True,
- line_width=0,
- )
-
- # Set y-axis ticks to show task names
- if line_plot._yticks is None:
- ax.set_ticks("y", list(y_positions), tasks)
-
- # Add legend if requested
- if line_plot._legend and len(line_plot.line_data) > 0:
- ax.set_legend(position="north east")
-
- return fig
-
- def _get_tikzfigure_axis_dimensions(self) -> tuple[str | None, str | None]:
- if self._width is None:
- return None, None
-
- total_width_in, total_height_in = set_size(
- width=self._width,
- ratio=self._ratio,
- dpi=self._dpi if self._dpi is not None else 300,
- )
- total_width_cm = total_width_in * 2.54
- total_height_cm = total_height_in * 2.54
- horizontal_sep_cm = getattr(TikzFigure, "GROUPPLOT_HORIZONTAL_SEP_CM", 1.5)
- available_width_cm = total_width_cm - horizontal_sep_cm * (self.ncols - 1)
- if available_width_cm <= 0:
- raise ValueError(
- f'Canvas width "{self._width}" is too small for {self.ncols} '
- "tikzfigure subplot(s)."
+ if verbose:
+ print("Drawing the canvas with Matplotlib for the tikzfigure backend")
+ # drawing with Matplotlib changes the global style and this canvas's
+ # record of its Matplotlib figure; neither is meant to change here
+ state = {
+ name: getattr(self, name)
+ for name in (
+ "_plotted",
+ "_matplotlib_fig",
+ "_matplotlib_axes",
+ "_matplotlib_twin_axes",
+ "_matplotlib_twiny_axes",
)
-
- axis_width_cm = available_width_cm / self.ncols
- return f"{axis_width_cm:.6g}cm", f"{total_height_cm:.6g}cm"
+ if hasattr(self, name)
+ }
+ with plt.rc_context(), plt.ioff():
+ fig, _ = self.plot_matplotlib(savefig=False, layers=layers, verbose=verbose)
+ try:
+ tikz = figure_to_tikz(
+ fig,
+ raster_dpi=raster_dpi,
+ max_markers=max_markers,
+ max_items=max_items,
+ max_points=max_points,
+ precision=precision,
+ )
+ finally:
+ plt.close(fig)
+ for name in (
+ "_plotted",
+ "_matplotlib_fig",
+ "_matplotlib_axes",
+ "_matplotlib_twin_axes",
+ "_matplotlib_twiny_axes",
+ ):
+ if name in state:
+ setattr(self, name, state[name])
+ elif hasattr(self, name):
+ delattr(self, name)
+ if verbose:
+ print(f"Converted {len(tikz.axes)} axes")
+ return tikz
def plot_plotext(
self,
@@ -2687,6 +2789,7 @@ def plot_plotext(
layers: list | None = None,
verbose: bool = False,
) -> PlotextFigure:
+ self._validate_import_backend("plotext")
if self._twinx_subplots:
raise NotImplementedError(
"twinx plots are not supported by the plotext backend"
@@ -2734,6 +2837,7 @@ def plot_plotly(
"""
+ self._validate_import_backend("plotly", allow_unsupported=allow_unsupported)
resolved_usetex = self._usetex if usetex is None else usetex
for subplot in self._subplot_dict.values():
@@ -2771,12 +2875,30 @@ def plot_plotly(
subplot_titles=subplot_titles,
specs=specs,
)
+ if self._hide_empty_subplots:
+ for row in range(self.nrows):
+ for col in range(self.ncols):
+ if (row, col) not in self._subplot_dict:
+ fig.update_xaxes(visible=False, row=row + 1, col=col + 1)
+ fig.update_yaxes(visible=False, row=row + 1, col=col + 1)
# Plot each subplot and propagate axis labels/scale
+ legend_names = set()
for (row, col), line_plot in self._subplot_dict.items():
traces, shapes, annotations = line_plot.plot_plotly(
layers=layers, allow_unsupported=allow_unsupported
)
+ if self._facet:
+ # One legend entry per label across all subplots; clicking it
+ # toggles the matching trace in every subplot.
+ for trace in traces:
+ name = getattr(trace, "name", None)
+ if not name or trace.type in ("pie", "table"):
+ continue
+ trace.legendgroup = name
+ if name in legend_names:
+ trace.showlegend = False
+ legend_names.add(name)
for trace in traces:
if trace.type in ("pie", "table"):
fig.add_trace(trace)
@@ -3037,11 +3159,13 @@ def plot_plotly(
self._plotly_fig = fig
return fig
- def _save_plotly(self, fig, filename: str) -> None:
+ def _save_plotly(
+ self, fig, filename: str, include_plotlyjs: bool | str = True
+ ) -> None:
_, extension = os.path.splitext(filename)
extension = extension.lower()
if extension in {".html", ".htm"}:
- fig.write_html(filename)
+ fig.write_html(filename, include_plotlyjs=include_plotlyjs)
return
try:
fig.write_image(filename)
@@ -3051,6 +3175,33 @@ def _save_plotly(self, fig, filename: str) -> None:
"(e.g., `pip install -U kaleido`), or export to HTML instead."
) from exc
+ def to_html(
+ self,
+ layers: list | None = None,
+ include_plotlyjs: bool | str = "cdn",
+ full_html: bool = True,
+ verbose: bool = False,
+ allow_unsupported: bool = False,
+ ) -> str:
+ """Render the canvas with the Plotly backend and return standalone HTML.
+
+ The returned markup embeds the figure data and calls plotly.js to
+ render it in a browser. By default ``include_plotlyjs="cdn"``
+ references plotly.js from a CDN instead of bundling it, producing a
+ much smaller string suited to embedding in an existing page; pass
+ ``include_plotlyjs=True`` to inline plotly.js for offline use, as
+ ``savefig(..., backend="plotly")`` does. Set ``full_html=False`` to
+ get just the ````/``