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") ![](README_files/figure-commonmark/cell-14-output-1.png) -### 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() +``` + +![](README_files/figure-commonmark/cell-20-output-1.png) + + (
, + 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 ``
``/``