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..8f4ffa3 100644
--- a/.github/workflows/docs.yml
+++ b/.github/workflows/docs.yml
@@ -33,6 +33,8 @@ jobs:
- 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
index ab33b36..cde5a4c 100644
--- a/.github/workflows/matplotlib-import.yml
+++ b/.github/workflows/matplotlib-import.yml
@@ -19,5 +19,7 @@ jobs:
- 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 02dfd59..11f4beb 100644
--- a/README.md
+++ b/README.md
@@ -219,23 +219,29 @@ canvas.show(backend="tikzfigure")

-### Horizontal Subplots with TikZ Backend
+### Subplots and Meshes with the TikZ Backend
-The tikzfigure backend supports creating side-by-side subplots (1×n
-layouts):
+The tikzfigure backend draws the canvas with Matplotlib and converts the
+drawn figure into pgfplots axes, so every layout converts (rows, columns,
+grids, twin axes), with LaTeX text, legends and colorbars. Lines,
+markers, bars and text are pgfplots code; meshes and images are included
+as images:
``` python
x = np.linspace(0, 2 * np.pi, 200)
-canvas, (ax1, ax2) = Canvas.subplots(ncols=2, width="10cm", ratio=0.3)
+canvas, (ax1, ax2) = Canvas.subplots(ncols=2, width="12cm", ratio=0.45)
-ax1.plot(x, np.sin(x), color="royalblue")
-ax1.set_title("sin(x)")
+ax1.plot(x, np.sin(x), color="royalblue", label="$\\sin x$")
+ax1.plot(x, np.cos(x), color="tomato", label="$\\cos x$")
+ax1.set_title("Lines")
+ax1.set_legend(True)
-ax2.plot(x, np.cos(x), color="tomato")
-ax2.set_title("cos(x)")
+xx, yy = np.meshgrid(x, x)
+ax2.pcolormesh(xx, yy, np.sin(xx) * np.cos(yy), cmap="RdBu_r")
+ax2.add_colorbar(label="$\\sin x \\cos y$")
+ax2.set_title("A mesh")
-canvas.suptitle("Trigonometric Functions")
-canvas.show(backend="tikzfigure") # Generates LaTeX subfigures
+canvas.show(backend="tikzfigure") # compiles with pdflatex
```
@@ -248,9 +254,11 @@ Figure 2
-**Note:** Only horizontal layouts (1×n) are currently supported with the
-tikzfigure backend. Vertical/grid layouts will raise
-`NotImplementedError`. See the tutorials for more examples.
+`canvas.render(backend="tikzfigure").savefig("figure.tikz")` writes the
+code for `\\input` in a LaTeX document, with the images next to it. Any
+Matplotlib figure converts the same way with
+`maxplotlib.backends.tikzfigure.figure_to_tikz(fig)`. See the tutorials
+for more examples.
### Terminal Backend with plotext
diff --git a/README.qmd b/README.qmd
index cefdf37..13111d0 100644
--- a/README.qmd
+++ b/README.qmd
@@ -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
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/pyproject.toml b/pyproject.toml
index 2ca1936..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,7 +19,7 @@ dependencies = [
"pint",
"plotly",
"plotext >= 6.0, < 7",
- "tikzfigure[vis]>=0.3.0",
+ "tikzfigure[vis]>=0.4.0",
]
[project.optional-dependencies]
test = [
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 ee66087..fb08d64 100644
--- a/src/maxplotlib/canvas/canvas.py
+++ b/src/maxplotlib/canvas/canvas.py
@@ -18,13 +18,7 @@
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
@@ -2250,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}")
@@ -2708,332 +2702,86 @@ 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.
"""
- self._validate_import_backend("tikzfigure")
- 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."
- )
+ from maxplotlib.backends.tikzfigure import figure_to_tikz
- 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,
- )
-
- # 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_plot._shift_x(line_data["x"])
- y = line_plot._shift_y(line_data["y"])
- 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_plot._shift_x(line_data["x"])
- y = line_plot._shift_y(line_data["y"])
- 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,
diff --git a/src/maxplotlib/subfigure/line_plot.py b/src/maxplotlib/subfigure/line_plot.py
index 3107353..da86d4f 100644
--- a/src/maxplotlib/subfigure/line_plot.py
+++ b/src/maxplotlib/subfigure/line_plot.py
@@ -3,7 +3,6 @@
import matplotlib.pyplot as plt
import numpy as np
import plotly.graph_objects as go
-from tikzfigure import TikzFigure
from maxplotlib.utils import xarray_support
@@ -21,26 +20,6 @@ def _mpl_mappable_kwargs(line):
return {k: v for k, v in line["kwargs"].items() if k != "colorbar"}
-_TIKZ_SUPPORTED_PLOT_TYPES = {
- "plot",
- "scatter",
- "bar",
- "barh",
- "fill_between",
- "errorbar",
- "step",
- "stairs",
- "stem",
- "hlines",
- "vlines",
- "axvspan",
- "axhspan",
- "fill",
- "gantt",
- "flame_chart",
-}
-
-
def _sample_colormap(colormap, count, *, css=True):
"""Sample ``count`` colors from a colormap, for either backend.
@@ -133,64 +112,6 @@ def _colormap_to_plotly_colorscale(colormap, steps=17):
return [[float(position), color] for position, color in zip(positions, colors)]
-def _tikz_style_kwargs(kwargs, *, default_color="black"):
- """Translate common Matplotlib-style options to pgfplots/TikZ options."""
- kwargs = dict(kwargs)
- style = {}
- if kwargs.get("color") is not None:
- style["color"] = kwargs["color"]
- else:
- style["color"] = default_color
- if kwargs.get("linewidth") is not None:
- style["line_width"] = kwargs["linewidth"]
- if kwargs.get("alpha") is not None:
- style["opacity"] = kwargs["alpha"]
- if kwargs.get("linestyle") in {"--", "dashed"}:
- style["dash_pattern"] = "on 4pt off 2pt"
- elif kwargs.get("linestyle") in {":", "dotted"}:
- style["dash_pattern"] = "on 1pt off 2pt"
- elif kwargs.get("linestyle") == "-.":
- style["dash_pattern"] = "on 4pt off 2pt on 1pt off 2pt"
- if kwargs.get("marker") is not None:
- style["mark"] = kwargs["marker"]
- if kwargs.get("markersize") is not None:
- style["mark_size"] = f"{kwargs['markersize']}pt"
- return style
-
-
-def _tikz_error_bounds(error, values):
- """Return lower and upper error arrays in Matplotlib's common formats."""
- if error is None:
- return None
- error = np.asarray(error, dtype=float)
- values = np.asarray(values, dtype=float)
- if error.ndim == 0:
- error = np.full(values.shape, error.item())
- if error.ndim == 2 and error.shape[0] == 2:
- return error[0], error[1]
- return error, error
-
-
-def _tikz_step_coordinates(x, y, where="pre"):
- """Expand line data into explicit coordinates for a stepped path."""
- x = np.asarray(x)
- y = np.asarray(y)
- if len(x) < 2:
- return x, y
- if where == "post":
- step_x = np.repeat(x, 2)[1:]
- step_y = np.repeat(y, 2)[:-1]
- elif where == "mid":
- mids = (x[:-1] + x[1:]) / 2
- step_x = np.ravel(np.column_stack((x[:-1], mids, mids, x[1:])))
- step_y = np.ravel(np.column_stack((y[:-1], y[:-1], y[1:], y[1:])))
- return step_x, step_y
- else:
- step_x = np.repeat(x, 2)[:-1]
- step_y = np.repeat(y, 2)[1:]
- return step_x, step_y
-
-
class Node:
def __init__(self, x, y, label="", content="", layer=0, **kwargs):
self.x = x
@@ -365,6 +286,9 @@ def _add(self, obj, layer):
for key in _NEUTRAL_KWARGS:
if key in kwargs:
obj[key] = kwargs.pop(key)
+ # TikZ's spelling of Matplotlib's linewidth, accepted by every backend
+ if "line_width" in kwargs and "linewidth" not in kwargs:
+ kwargs["linewidth"] = kwargs.pop("line_width")
for key in _NEUTRAL_KWARGS:
obj.setdefault(key, None)
self.line_data.append(obj)
@@ -1799,6 +1723,13 @@ def plot_matplotlib(
)
ax.set_ylim(-0.5, max_depth)
+ # patches do not autoscale the view; the subplot's own
+ # limits, if any, are applied afterwards
+ if n:
+ ax.set_xlim(
+ float(np.min(start_times)),
+ float(np.max(start_times + values)),
+ )
ax.set_ylabel("Stack Depth")
elif line["plot_type"] == "fill_between":
ax.fill_between(
@@ -2113,278 +2044,6 @@ def _tag_matplotlib_artists(ax, artists_before, meta):
except AttributeError:
continue
- def plot_tikzfigure(self, layers=None, verbose: bool = False) -> TikzFigure:
-
- tikz_figure = TikzFigure()
- for layer_name, layer_lines in self.layered_line_data.items():
- if layers and layer_name not in layers:
- continue
- for line in layer_lines:
- plot_type = line["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":
- x = self._shift_x(line["x"])
- y = self._shift_y(line["y"])
-
- nodes = [[xi, yi] for xi, yi in zip(x, y)]
- tikz_figure.draw(
- nodes=nodes,
- **_tikz_style_kwargs(line["kwargs"]),
- )
- elif plot_type == "scatter":
- x = self._shift_x(line["x"])
- y = self._shift_y(line["y"])
- style = _tikz_style_kwargs(line["kwargs"])
- style.setdefault("mark", "*")
- style["line_width"] = 0
- tikz_figure.draw(
- nodes=[[xi, yi] for xi, yi in zip(x, y)],
- **style,
- )
- elif plot_type in {"bar", "barh"}:
- kwargs = line["kwargs"]
- style = _tikz_style_kwargs(kwargs)
- style["fill"] = kwargs.get("color", "blue")
- style["fill_opacity"] = kwargs.get("alpha", 1.0)
- style["line_width"] = kwargs.get("linewidth", 0)
- if plot_type == "bar":
- width = kwargs.get("width", 0.8)
- for x, height in zip(line["x"], line["height"]):
- x = self._shift_x(x)
- height = height * self._yscale
- tikz_figure.draw(
- nodes=[
- [x - width / 2, 0],
- [x + width / 2, 0],
- [x + width / 2, height],
- [x - width / 2, height],
- ],
- cycle=True,
- **style,
- )
- else:
- height = kwargs.get("height", 0.8)
- for y, width in zip(line["y"], line["width"]):
- y = self._shift_y(y)
- width = width * self._xscale
- tikz_figure.draw(
- nodes=[
- [0, y - height / 2],
- [width, y - height / 2],
- [width, y + height / 2],
- [0, y + height / 2],
- ],
- cycle=True,
- **style,
- )
- elif plot_type == "fill_between":
- x = self._shift_x(line["x"])
- y1 = np.asarray(line["y1"])
- y2 = np.broadcast_to(line["y2"], y1.shape)
- nodes = [[xi, yi] for xi, yi in zip(x, y1)]
- nodes.extend([[xi, yi] for xi, yi in zip(x[::-1], y2[::-1])])
- kwargs = line["kwargs"]
- style = _tikz_style_kwargs(kwargs)
- style["fill"] = kwargs.get("color", "blue")
- style["fill_opacity"] = kwargs.get("alpha", 0.25)
- tikz_figure.draw(nodes=nodes, cycle=True, **style)
- elif plot_type == "errorbar":
- x = self._shift_x(line["x"])
- y = self._shift_y(line["y"])
- style = _tikz_style_kwargs(line["kwargs"])
- tikz_figure.draw(nodes=[[xi, yi] for xi, yi in zip(x, y)], **style)
- y_bounds = _tikz_error_bounds(line["yerr"], y)
- if y_bounds is not None:
- lower, upper = y_bounds
- for xi, low, high in zip(x, y - lower, y + upper):
- tikz_figure.draw(nodes=[[xi, low], [xi, high]], **style)
- x_bounds = _tikz_error_bounds(line["xerr"], x)
- if x_bounds is not None:
- lower, upper = x_bounds
- for yi, low, high in zip(y, x - lower, x + upper):
- tikz_figure.draw(nodes=[[low, yi], [high, yi]], **style)
- elif plot_type in {"step", "stairs"}:
- kwargs = line["kwargs"]
- if plot_type == "step":
- x = line["x"]
- y = line["y"]
- where = kwargs.get("where", "pre")
- else:
- values = line["values"]
- edges = line["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)
- x = self._shift_x(x)
- y = self._shift_y(y)
- tikz_figure.draw(
- nodes=[[xi, yi] for xi, yi in zip(x, y)],
- **_tikz_style_kwargs(kwargs),
- )
- elif plot_type == "stem":
- x = self._shift_x(line["x"])
- y = self._shift_y(line["y"])
- kwargs = line["kwargs"]
- style = _tikz_style_kwargs(kwargs)
- marker_style = dict(style)
- marker_style.update(mark=kwargs.get("marker", "*"), line_width=0)
- tikz_figure.draw(
- nodes=[[xi, yi] for xi, yi in zip(x, y)], **marker_style
- )
- for xi, yi in zip(x, y):
- tikz_figure.draw(nodes=[[xi, 0], [xi, yi]], **style)
- elif plot_type in {"hlines", "vlines"}:
- kwargs = _tikz_style_kwargs(line["kwargs"])
- if plot_type == "hlines":
- for yi, left, right in zip(
- np.atleast_1d(line["y"]),
- np.atleast_1d(line["xmin"]),
- np.atleast_1d(line["xmax"]),
- ):
- tikz_figure.draw(nodes=[[left, yi], [right, yi]], **kwargs)
- else:
- for xi, bottom, top in zip(
- np.atleast_1d(line["x"]),
- np.atleast_1d(line["ymin"]),
- np.atleast_1d(line["ymax"]),
- ):
- tikz_figure.draw(nodes=[[xi, bottom], [xi, top]], **kwargs)
- elif plot_type in {"axvspan", "axhspan"}:
- kwargs = line["kwargs"]
- style = _tikz_style_kwargs(kwargs)
- style["fill"] = kwargs.get("color", "blue")
- style["fill_opacity"] = kwargs.get("alpha", 0.2)
- if plot_type == "axvspan":
- ymin, ymax = self._ymin or 0, self._ymax or 1
- nodes = [
- [line["xmin"], ymin],
- [line["xmax"], ymin],
- [line["xmax"], ymax],
- [line["xmin"], ymax],
- ]
- else:
- xmin, xmax = self._xmin or 0, self._xmax or 1
- nodes = [
- [xmin, line["ymin"]],
- [xmax, line["ymin"]],
- [xmax, line["ymax"]],
- [xmin, line["ymax"]],
- ]
- tikz_figure.draw(nodes=nodes, cycle=True, **style)
- elif plot_type == "fill":
- if len(line["args"]) < 2:
- raise ValueError("tikzfigure fill requires x and y coordinates")
- x, y = line["args"][:2]
- kwargs = line["kwargs"]
- style = _tikz_style_kwargs(kwargs)
- style["fill"] = kwargs.get("color", "blue")
- style["fill_opacity"] = kwargs.get("alpha", 0.25)
- tikz_figure.draw(
- nodes=[[xi, yi] for xi, yi in zip(x, y)],
- cycle=True,
- **style,
- )
- elif line["plot_type"] == "gantt":
- tasks = line["tasks"]
- start_times = self._shift_x(line["start_times"])
- durations = line["durations"] * self._xscale
- y_positions = np.arange(len(tasks))
-
- # Draw horizontal bars for each task
- for i, (task, start, duration) in enumerate(
- zip(tasks, start_times, durations)
- ):
- # Create rectangle nodes for the bar
- x_start = start
- x_end = start + duration
- y_pos = y_positions[i]
- bar_height = 0.8 # Bar thickness
-
- # Draw rectangle as a path
- rect_nodes = [
- [x_start, y_pos - bar_height / 2],
- [x_end, y_pos - bar_height / 2],
- [x_end, y_pos + bar_height / 2],
- [x_start, y_pos + bar_height / 2],
- ]
- tikz_figure.draw(
- nodes=rect_nodes,
- cycle=True,
- fill=line["kwargs"].get("color", "blue"),
- **line["kwargs"],
- )
- elif line["plot_type"] == "flame_chart":
- labels = line["labels"]
- parents = line["parents"]
- values = line["values"] * self._xscale
- start_times = line["start_times"]
-
- # Calculate depths
- n = len(labels)
- depths = np.zeros(n, dtype=int)
- if start_times is None:
- start_times = np.zeros(n)
- else:
- start_times = self._shift_x(start_times)
-
- for i in range(n):
- if parents[i] is None:
- depths[i] = 0
- else:
- parent_idx = (
- parents[i]
- if isinstance(parents[i], int)
- else list(labels).index(parents[i])
- )
- depths[i] = depths[parent_idx] + 1
-
- # Draw rectangles for each frame
- bar_height = 0.8
- explicit_colors = line["kwargs"].get("colors")
- if isinstance(explicit_colors, str) or not hasattr(
- explicit_colors, "__len__"
- ):
- explicit_colors = (
- None if explicit_colors is None else [explicit_colors]
- )
- colors = ["red", "blue", "green", "orange", "purple", "cyan"]
-
- for i in range(n):
- x_start = start_times[i]
- x_end = start_times[i] + values[i]
- y_pos = depths[i]
- if explicit_colors:
- color = explicit_colors[i % len(explicit_colors)]
- else:
- color = colors[depths[i] % len(colors)]
-
- rect_nodes = [
- [x_start, y_pos - bar_height / 2],
- [x_end, y_pos - bar_height / 2],
- [x_end, y_pos + bar_height / 2],
- [x_start, y_pos + bar_height / 2],
- ]
- tikz_figure.draw(
- nodes=rect_nodes,
- cycle=True,
- fill=color,
- **{
- k: v
- for k, v in line["kwargs"].items()
- if k not in ("colormap", "colors")
- },
- )
- if verbose:
- print("Generated TikZ figure:")
- print(tikz_figure.generate_tikz())
- return tikz_figure
-
def plot_plotly(self, layers=None, allow_unsupported=False):
if hasattr(self, "_import_projection"):
raise NotImplementedError(
diff --git a/src/maxplotlib/tests/test_canvas.py b/src/maxplotlib/tests/test_canvas.py
index e0622a3..9d6922d 100644
--- a/src/maxplotlib/tests/test_canvas.py
+++ b/src/maxplotlib/tests/test_canvas.py
@@ -1,3 +1,6 @@
+import re
+
+
def test():
pass
@@ -68,30 +71,37 @@ def test_canvas_plot_tikzfigure_respects_width_and_ratio():
tikz = canvas.plot_tikzfigure().generate_tikz()
- assert "width=10cm" in tikz
- assert "height=20cm" in tikz
+ # the axis box is the Matplotlib axes of the 10cm x 20cm figure
+ width = float(re.search(r"width=([0-9.]+)in", tikz).group(1))
+ height = float(re.search(r"height=([0-9.]+)in", tikz).group(1))
+ assert 0.5 * 10 / 2.54 < width < 10 / 2.54
+ assert 0.5 * 20 / 2.54 < height < 20 / 2.54
assert "title=Parabola" in tikz
-def test_canvas_plot_tikzfigure_vertical_not_supported():
- """Test that vertical layouts raise NotImplementedError."""
+def test_canvas_plot_tikzfigure_vertical_layout():
+ """A 2x1 layout gives two axes, the first above the second."""
import numpy as np
- import pytest
from maxplotlib import Canvas
x = np.linspace(0, 2 * np.pi, 50)
- # Create 2×1 layout (nrows=2)
canvas, axes = Canvas.subplots(nrows=2, width="10cm")
axes[0].plot(x, np.sin(x))
axes[1].plot(x, np.cos(x))
- # Should raise NotImplementedError
- with pytest.raises(NotImplementedError) as exc_info:
- canvas.plot_tikzfigure()
+ figure = canvas.plot_tikzfigure()
+ tikz = figure.generate_tikz()
- assert "nrows > 1" in str(exc_info.value)
+ assert len(figure.axes) == 2
+ positions = [
+ (float(x), float(y))
+ for x, y in re.findall(r"at=\{\(([0-9.]+)in,([0-9.]+)in\)\}", tikz)
+ ]
+ assert len(positions) == 2
+ assert positions[0][1] > positions[1][1]
+ assert positions[0][0] == positions[1][0]
def test_tikzfigure_supports_scatter_bars_fills_and_errorbars():
@@ -109,22 +119,27 @@ def test_tikzfigure_supports_scatter_bars_fills_and_errorbars():
tikz = canvas.render(backend="tikzfigure").generate_tikz()
assert "mark=*" in tikz
- assert "fill=blue" in tikz
- assert "fill=green" in tikz
+ assert "\\definecolor{mpl0000FF}{HTML}{0000FF}" in tikz
+ assert "fill=mpl0000FF" in tikz
+ assert "fill=mpl008000, fill opacity=0.2" in tikz
assert tikz.count("coordinates") >= 4
-def test_tikzfigure_rejects_unsupported_plot_types_explicitly():
+def test_tikzfigure_draws_images_as_graphics():
import numpy as np
- import pytest
from maxplotlib import Canvas
canvas = Canvas()
- canvas.imshow(np.ones((2, 2)))
+ canvas.imshow(np.arange(4.0).reshape(2, 2))
+
+ figure = canvas.render(backend="tikzfigure")
+ tikz = figure.generate_tikz()
- with pytest.raises(NotImplementedError, match="imshow"):
- canvas.render(backend="tikzfigure")
+ assert "\\addplot[forget plot] graphics" in tikz
+ ((name, data),) = figure.files().items()
+ assert name in tikz
+ assert data.startswith(b"\x89PNG")
def test_tikzfigure_supports_step_stem_reference_lines_spans_and_fill():
@@ -135,7 +150,7 @@ def test_tikzfigure_supports_step_stem_reference_lines_spans_and_fill():
x = np.arange(4)
canvas = Canvas()
canvas.step(x, [1, 2, 1, 3], color="black")
- canvas.stem(x, [1, 2, 1, 3], color="purple")
+ canvas.stem(x, [1, 2, 1, 3], linefmt="C4-", markerfmt="C4o")
canvas.hlines([1, 2], 0, 3, color="gray")
canvas.vlines([1, 2], 0, 3, color="gray")
canvas.axvspan(1, 2, color="orange", alpha=0.2)
@@ -145,9 +160,9 @@ def test_tikzfigure_supports_step_stem_reference_lines_spans_and_fill():
tikz = canvas.render(backend="tikzfigure").generate_tikz()
assert "mark=*" in tikz
- assert "fill=orange" in tikz
- assert "fill=cyan" in tikz
- assert tikz.count("coordinates") >= 10
+ assert "fill=mplFFA500" in tikz
+ assert "fill=mpl00FFFF" in tikz
+ assert tikz.count("coordinates") >= 7
def test_canvas_matplotlib_gridspec_kw_affects_row_spacing():
diff --git a/src/maxplotlib/tests/test_flame_chart.py b/src/maxplotlib/tests/test_flame_chart.py
index 64c5e91..4c1b4a5 100644
--- a/src/maxplotlib/tests/test_flame_chart.py
+++ b/src/maxplotlib/tests/test_flame_chart.py
@@ -2,6 +2,8 @@
Tests for flame chart functionality across all backends.
"""
+import shutil
+
import numpy as np
import pytest
@@ -162,6 +164,13 @@ def test_flame_chart_tikzfigure_backend(sample_flame_data, tmp_path):
canvas.set_ylabel("Stack Depth")
canvas.set_title("Test Flame Chart")
+ # the TikZ code needs no LaTeX; compiling it to PDF needs pdflatex
+ tikz_file = tmp_path / "test_flame_tikz.tikz"
+ canvas.savefig(str(tikz_file), backend="tikzfigure")
+ assert "\\begin{axis}" in tikz_file.read_text()
+
+ if shutil.which("pdflatex") is None:
+ pytest.skip("pdflatex not installed")
output_file = tmp_path / "test_flame_tikz.pdf"
canvas.savefig(str(output_file), backend="tikzfigure")
diff --git a/src/maxplotlib/tests/test_gantt_chart.py b/src/maxplotlib/tests/test_gantt_chart.py
index 66e24fd..496e77e 100644
--- a/src/maxplotlib/tests/test_gantt_chart.py
+++ b/src/maxplotlib/tests/test_gantt_chart.py
@@ -2,6 +2,8 @@
Tests for gantt chart functionality across all backends.
"""
+import shutil
+
import numpy as np
import pytest
@@ -133,6 +135,13 @@ def test_gantt_chart_tikzfigure_backend(sample_gantt_data, tmp_path):
canvas.set_xlabel("Time (days)")
canvas.set_title("Project Timeline")
+ # the TikZ code needs no LaTeX; compiling it to PDF needs pdflatex
+ tikz_file = tmp_path / "test_gantt_tikz.tikz"
+ canvas.savefig(str(tikz_file), backend="tikzfigure")
+ assert "\\begin{axis}" in tikz_file.read_text()
+
+ if shutil.which("pdflatex") is None:
+ pytest.skip("pdflatex not installed")
output_file = tmp_path / "test_gantt_tikz.pdf"
canvas.savefig(str(output_file), backend="tikzfigure")
diff --git a/src/maxplotlib/tests/test_tikzfigure_backend.py b/src/maxplotlib/tests/test_tikzfigure_backend.py
new file mode 100644
index 0000000..a496fb8
--- /dev/null
+++ b/src/maxplotlib/tests/test_tikzfigure_backend.py
@@ -0,0 +1,259 @@
+"""The tikzfigure backend: drawn Matplotlib figures as pgfplots axes."""
+
+import re
+import shutil
+import warnings
+
+import matplotlib
+
+matplotlib.use("Agg")
+import matplotlib.pyplot as plt # noqa: E402
+import numpy as np # noqa: E402
+import pytest # noqa: E402
+
+from maxplotlib import Canvas # noqa: E402
+from maxplotlib.backends.tikzfigure import ( # noqa: E402
+ TikzConversionWarning,
+ figure_to_tikz,
+ latex,
+)
+
+needs_pdflatex = pytest.mark.skipif(
+ shutil.which("pdflatex") is None, reason="pdflatex not installed"
+)
+
+
+@pytest.fixture(autouse=True)
+def close_figures():
+ yield
+ plt.close("all")
+
+
+def axes_options(tikz):
+ """The option lists of every ``axis`` environment."""
+ return re.findall(r"\\begin\{axis\}\[(.*?)\]\n", tikz)
+
+
+# -- text ------------------------------------------------------------------
+@pytest.mark.parametrize(
+ "text, expected",
+ [
+ ("plain", "plain"),
+ ("50% of a_b & #1", r"50\% of a\_b \& \#1"),
+ (r"$\gamma/\omega_{ci}$ fit", r"$\gamma/\omega_{ci}$ fit"),
+ (r"$\mathdefault{10^{-3}}$", "${10^{-3}}$"),
+ (r"$\mathdefault{0.5}$", "$0.5$"),
+ ("ω = 2", r"$\omega$ = 2"),
+ (r"$ω_p$", r"$\omega _p$"),
+ ("−1", "$-$1"),
+ ("cost: $5", r"cost: \$5"),
+ ("two\nlines", r"two\\lines"),
+ ("", ""),
+ (None, ""),
+ ],
+)
+def test_latex(text, expected):
+ assert latex(text) == expected
+
+
+# -- what is drawn ---------------------------------------------------------
+def test_lines_markers_and_styles():
+ fig, ax = plt.subplots()
+ ax.plot([0, 1, 2], [0, 1, 0], "--", color="red", lw=2)
+ ax.plot([0, 1, 2], [1, 2, 1], "s", ms=6, mfc="none", mec="blue")
+ tikz = figure_to_tikz(fig).generate_tikz()
+ assert "\\definecolor{mplFF0000}{HTML}{FF0000}" in tikz
+ assert "draw=mplFF0000, line width=2pt, dash pattern=on 7.4pt off 3.2pt" in tikz
+ assert "only marks, mark=square, mark size=3pt" in tikz
+ assert "fill=none, draw=mpl0000FF" in tikz
+
+
+def test_the_legend_keeps_matplotlib_order_and_style():
+ fig, ax = plt.subplots()
+ ax.fill_between([0, 1], [0, 1], alpha=0.5, label="area")
+ ax.plot([0, 1], [1, 0], color="k", label="line")
+ ax.legend(handles=ax.lines + ax.collections, loc="upper left")
+ tikz = figure_to_tikz(fig).generate_tikz()
+ entries = re.findall(r"\\addlegendentry\{(.*?)\}", tikz)
+ assert entries == ["line", "area"]
+ images = re.findall(r"\\addlegendimage\{(.*?)\}\n", tikz)
+ assert "area legend" in images[1] and "draw=mpl000000" in images[0]
+ assert "legend style={at={(" in tikz and "anchor=north west" in tikz
+ assert tikz.count("forget plot") >= 2
+
+
+def test_log_and_reversed_axes():
+ fig, ax = plt.subplots()
+ ax.semilogy([1, 2, 3], [1, 10, 100])
+ ax.invert_xaxis()
+ (options,) = axes_options(figure_to_tikz(fig).generate_tikz())
+ assert "ymode=log" in options
+ assert "x dir=reverse" in options
+ assert "ytick={1,10,100}" in options
+
+
+def test_text_and_annotation_arrows_are_not_clipped():
+ fig, ax = plt.subplots()
+ ax.plot([0, 1], [0, 1])
+ ax.set_xlim(0, 1)
+ ax.set_ylim(0, 1)
+ ax.annotate(
+ "peak, 50%", xy=(0.5, 0.5), xytext=(0.8, 1.05), arrowprops={"arrowstyle": "->"}
+ )
+ ax.text(0.1, 0.9, "$x_0$", ha="left", va="top")
+ tikz = figure_to_tikz(fig).generate_tikz()
+ assert r"{peak, 50\%}" in tikz
+ assert "anchor=north west" in tikz and "{$x_0$}" in tikz
+ assert re.search(r"\\draw\[.*\] \(axis cs:", tikz), "the arrow is a \\draw"
+
+
+def test_far_away_geometry_is_clipped_and_text_outside_the_figure_left_out():
+ fig, ax = plt.subplots()
+ ax.plot([0, 1e6], [0, 1])
+ ax.text(1e6, 0.5, "far away")
+ ax.set_xlim(0, 1)
+ tikz = figure_to_tikz(fig).generate_tikz()
+ xs = [float(x) for x in re.findall(r"\(([-0-9.e+]+),[-0-9.e+]+\)", tikz)]
+ assert xs and max(xs) < 3
+ assert "far away" not in tikz
+
+
+def test_scatter_colored_by_value_is_one_plot_per_color():
+ fig, ax = plt.subplots()
+ ax.scatter([0, 1, 2, 3], [0, 1, 2, 3], c=[0, 0, 1, 1], s=20, cmap="viridis")
+ tikz = figure_to_tikz(fig).generate_tikz()
+ assert tikz.count("only marks, mark=*") == 2
+
+
+def test_large_scatter_and_meshes_are_images(tmp_path):
+ fig, ax = plt.subplots()
+ mesh = ax.pcolormesh(np.random.default_rng(0).random((10, 20)))
+ fig.colorbar(mesh, label="value [a.u.]")
+ ax.scatter(*np.random.default_rng(1).random((2, 50)), s=2)
+ tikz_figure = figure_to_tikz(fig, max_markers=10)
+ tikz = tikz_figure.generate_tikz()
+ assert len(tikz_figure.axes) == 2, "the axes and the colorbar"
+ assert tikz.count("\\addplot[forget plot] graphics") == 2
+ main, colorbar = axes_options(tikz)
+ assert "axis on top" in main
+ assert "ylabel={value [a.u.]}" in colorbar and "xtick=\\empty" in colorbar
+ tikz_figure.savefig(tmp_path / "figure.tikz")
+ images = sorted(path.name for path in tmp_path.glob("*.png"))
+ assert len(images) == 2
+ assert all(name in tikz for name in images)
+
+
+def test_twin_axes_share_the_position_with_ticks_on_the_right():
+ fig, ax = plt.subplots()
+ ax.plot([0, 1], [0, 1])
+ twin = ax.twinx()
+ twin.plot([0, 1], [1, 0], color="C1")
+ twin.set_ylabel("right")
+ first, second = axes_options(figure_to_tikz(fig).generate_tikz())
+ position = re.compile(r"at=\{\(([0-9.]+)in,([0-9.]+)in\)\}")
+ assert position.search(first).groups() == position.search(second).groups()
+ assert "ytick pos=right" in second and "ylabel near ticks" in second
+ assert "xtick=\\empty" in second
+
+
+def test_hidden_spines_and_tick_labels():
+ fig, axes = plt.subplots(2, 1, sharex=True)
+ for ax in axes:
+ ax.plot([0, 1], [0, 1])
+ ax.spines[["top", "right"]].set_visible(False)
+ upper, lower = axes_options(figure_to_tikz(fig).generate_tikz())
+ assert "axis x line*=bottom" in upper and "axis y line*=left" in upper
+ assert "xticklabels={}" in upper and "xticklabels={}" not in lower
+
+
+def test_categories_keep_their_labels():
+ fig, ax = plt.subplots()
+ ax.bar(["low", "mid_1", "high"], [1, 3, 2])
+ (options,) = axes_options(figure_to_tikz(fig).generate_tikz())
+ assert r"xticklabels={{low},{mid\_1},{high}}" in options
+
+
+def test_polar_axes_are_an_image_with_a_warning():
+ fig = plt.figure()
+ ax = fig.add_subplot(projection="polar")
+ ax.plot([0, 1, 2], [1, 2, 1])
+ with pytest.warns(TikzConversionWarning, match="polar"):
+ tikz = figure_to_tikz(fig).generate_tikz()
+ assert "hide axis" in tikz and "graphics" in tikz
+
+
+def test_the_figure_is_left_as_it_was():
+ fig, ax = plt.subplots(layout="constrained")
+ (line,) = ax.plot([0, 1], [0, 1])
+ mesh = ax.pcolormesh(np.ones((2, 2)))
+ engine = fig.get_layout_engine()
+ figure_to_tikz(fig)
+ assert fig.get_layout_engine() is engine
+ assert line.get_visible() and mesh.get_visible() and ax.xaxis.get_visible()
+ assert fig.patch.get_visible()
+
+
+def test_figure_texts_are_placed_in_inches():
+ fig, ax = plt.subplots(figsize=(4, 3))
+ ax.plot([0, 1], [0, 1])
+ fig.suptitle("All of it")
+ tikz = figure_to_tikz(fig).generate_tikz()
+ match = re.search(
+ r"\\node\[anchor=north.*\] at \(([0-9.]+)in,([0-9.]+)in\) \{All of it\}", tikz
+ )
+ assert match
+ assert float(match.group(1)) == pytest.approx(2.0, abs=0.01)
+
+
+# -- through a Canvas ------------------------------------------------------
+def test_canvas_with_meshes_colorbars_and_a_grid_of_subplots():
+ canvas, axes = Canvas.subplots(nrows=2, ncols=2)
+ x = np.linspace(0, 1, 20)
+ axes[0][0].plot(x, x**2, label="square")
+ axes[0][0].set_legend(True)
+ axes[0][1].pcolormesh(x, x, np.outer(x, x), cmap="magma")
+ axes[0][1].add_colorbar(label="z")
+ axes[1][0].scatter(x, x, color="C2")
+ axes[1][1].imshow(np.eye(3))
+ figure = canvas.render(backend="tikzfigure")
+ assert len(figure.axes) == 5
+ assert len(figure.files()) == 3 # the mesh, the image and the colorbar strip
+
+
+def test_an_imported_figure_with_a_colorbar_converts():
+ fig, ax = plt.subplots()
+ image = ax.imshow(np.arange(6.0).reshape(2, 3))
+ fig.colorbar(image)
+ canvas = Canvas.from_matplotlib(fig)
+ figure = canvas.render(backend="tikzfigure")
+ assert len(figure.axes) == 2
+
+
+def test_rendering_leaves_the_canvas_and_global_style_alone():
+ canvas = Canvas()
+ canvas.plot([0, 1], [0, 1])
+ style = dict(plt.rcParams)
+ canvas.render(backend="tikzfigure")
+ assert dict(plt.rcParams) == style
+ assert not getattr(canvas, "_plotted", False)
+
+
+# -- compiled ---------------------------------------------------------------
+@needs_pdflatex
+def test_a_figure_with_everything_compiles(tmp_path):
+ fig, axes = plt.subplots(1, 2, figsize=(7, 3), layout="constrained")
+ t = np.linspace(0, 10, 3000)
+ axes[0].plot(t, np.sin(t) * np.exp(0.2 * t), label=r"$\sin t\, e^{t/5}$")
+ axes[0].axvspan(2, 3, alpha=0.2, color="C1", label="window, fit")
+ axes[0].errorbar([1, 5], [1, 2], yerr=0.5, fmt="o", capsize=3)
+ axes[0].set_yscale("symlog") # not pgfplots: an image
+ axes[0].legend()
+ mesh = axes[1].pcolormesh(np.random.default_rng(0).random((30, 30)))
+ axes[1].contour(np.random.default_rng(0).random((30, 30)), levels=[0.5], colors="k")
+ fig.colorbar(mesh, ax=axes[1], label="$|B|$ [T]")
+ fig.suptitle("Everything_1 & more")
+ with warnings.catch_warnings():
+ warnings.simplefilter("ignore", TikzConversionWarning)
+ figure = figure_to_tikz(fig)
+ figure.savefig(tmp_path / "figure.pdf")
+ assert (tmp_path / "figure.pdf").stat().st_size > 1000
diff --git a/tutorials/tutorial_02.ipynb b/tutorials/tutorial_02.ipynb
index e809ee6..651e27e 100644
--- a/tutorials/tutorial_02.ipynb
+++ b/tutorials/tutorial_02.ipynb
@@ -82,7 +82,7 @@
" xlim=(0, 360),\n",
" ylim=(-1.5, 1.5),\n",
" grid=True,\n",
- " caption=\"Sine Function\",\n",
+ " title=\"Sine Function\",\n",
" width=0.45,\n",
")\n",
"ax1.add_plot(x=x, y=y1, label=\"sin(x)\", color=\"red\", line_width=\"1.5pt\")\n",
@@ -95,7 +95,7 @@
" xlim=(0, 360),\n",
" ylim=(-1.5, 1.5),\n",
" grid=True,\n",
- " caption=\"Cosine Function\",\n",
+ " title=\"Cosine Function\",\n",
" width=0.45,\n",
")\n",
"ax2.add_plot(x=x, y=y2, label=\"cos(x)\", color=\"blue\", line_width=\"1.5pt\")\n",
diff --git a/tutorials/tutorial_07_tikz.ipynb b/tutorials/tutorial_07_tikz.ipynb
index 6fdf89b..95f2db0 100644
--- a/tutorials/tutorial_07_tikz.ipynb
+++ b/tutorials/tutorial_07_tikz.ipynb
@@ -47,7 +47,9 @@
"---\n",
"## Part 1 — Canvas → TikZ\n",
"\n",
- "The fastest path: build a plot with the standard Canvas API, then pass `backend='tikzfigure'` to get a `TikzFigure` object back."
+ "The fastest path: build a plot with the standard Canvas API, then pass `backend=\"tikzfigure\"` to get a `TikzFigure` object back.\n",
+ "\n",
+ "The canvas is drawn with Matplotlib first, off screen, and the drawn figure is converted: every subplot becomes a pgfplots `axis` at the place and size it has in the Matplotlib figure, with its limits, scales, labels, ticks, legend and colorbar. Lines, markers, bars, fills and text become pgfplots code; meshes and images become images placed with `\\addplot graphics`. So every layout and every plot type converts, and the TikZ figure shows what Matplotlib shows."
]
},
{
@@ -68,13 +70,14 @@
"x = np.linspace(0, 2 * np.pi, 60)\n",
"\n",
"canvas = Canvas(width=\"10cm\", ratio=0.6)\n",
- "canvas.plot(x, np.sin(x), label=\"sin\", color=\"steelblue\", line_width=1.5)\n",
- "canvas.plot(x, np.cos(x), label=\"cos\", color=\"tomato\", line_width=1.2)\n",
- "canvas.set_xlabel(\"x\")\n",
- "canvas.set_ylabel(\"y\")\n",
+ "canvas.plot(x, np.sin(x), label=\"sin\", color=\"steelblue\", linewidth=1.5)\n",
+ "canvas.plot(x, np.cos(x), label=\"cos\", color=\"tomato\", linewidth=1.2)\n",
+ "canvas.set_xlabel(\"$x$\")\n",
+ "canvas.set_ylabel(\"$y$\")\n",
"canvas.set_title(\"Trigonometric functions\")\n",
+ "canvas.set_legend(True)\n",
"\n",
- "# backend='tikzfigure' returns a TikzFigure object\n",
+ "# backend=\"tikzfigure\" returns a TikzFigure object\n",
"tikz = canvas.render(backend=\"tikzfigure\")\n",
"print(type(tikz))"
]
@@ -86,8 +89,7 @@
"source": [
"### Plotly preview\n",
"\n",
- "Before exporting to TikZ, you can preview the same `Canvas` interactively in a notebook using the Plotly backend:\n",
- "\n"
+ "Before exporting to TikZ, you can preview the same `Canvas` interactively in a notebook using the Plotly backend:"
]
},
{
@@ -117,8 +119,8 @@
"source": [
"### 1.2 Inspecting the generated LaTeX\n",
"\n",
- "`str(tikz)` returns the raw LaTeX source string (and `generate_tikz()` remains available explicitly). \n",
- "Each data line becomes a `\\draw` command connecting coordinate pairs."
+ "`str(tikz)` returns the raw LaTeX source string (and `generate_tikz()` remains available explicitly).\n",
+ "The colors are defined once with `\\definecolor`; the subplot is an `axis` placed with `at=` and sized with `width=`/`height=` (`scale only axis`), each line an `\\addplot`. The legend repeats Matplotlib's legend, entry by entry (`\\addlegendimage`, `\\addlegendentry`)."
]
},
{
@@ -136,10 +138,9 @@
"id": "10",
"metadata": {},
"source": [
- "### 1.2.1 Checking explicit width and height\n",
+ "### 1.3 Size and layout\n",
"\n",
- "When you set both `width=` and `ratio=`, the TikZ export now writes explicit pgfplots dimensions.\n",
- "This is useful when you want a tall figure for a column-sized layout in LaTeX."
+ "The TikZ figure has the canvas's size (`width=` and `ratio=`, or `figsize=`), and every axis the box Matplotlib gives it, so a grid of subplots, shared axes and twin axes keep their layout. Text is typeset by LaTeX at Matplotlib's font sizes."
]
},
{
@@ -149,19 +150,17 @@
"metadata": {},
"outputs": [],
"source": [
- "canvas_ratio2, ax_ratio2 = Canvas.subplots(width=\"10cm\", ratio=2)\n",
- "ax_ratio2.plot(x, np.exp(-x / np.pi), color=\"purple\", line_width=1.5)\n",
- "ax_ratio2.set_title(\"ratio = 2 export\")\n",
- "\n",
- "tikz_ratio2 = canvas_ratio2.render(backend=\"tikzfigure\")\n",
- "ratio2_code = tikz_ratio2.generate_tikz()\n",
- "\n",
- "for line in ratio2_code.splitlines():\n",
- " if \"nextgroupplot\" in line:\n",
- " print(line.strip())\n",
- " break\n",
+ "canvas_grid, axes = Canvas.subplots(\n",
+ " nrows=2, ncols=2, width=\"12cm\", ratio=0.7, hspace=0.5\n",
+ ")\n",
+ "for index, ax in enumerate(np.ravel(axes)):\n",
+ " ax.plot(x, np.sin((index + 1) * x), linewidth=1.2)\n",
+ " ax.set_title(f\"$n = {index + 1}$\")\n",
"\n",
- "# Expected: width=10cm and height=20cm in the \\nextgroupplot options"
+ "tikz_grid = canvas_grid.render(backend=\"tikzfigure\")\n",
+ "for line in tikz_grid.generate_tikz().splitlines():\n",
+ " if \"\\\\begin{axis}\" in line:\n",
+ " print(line.strip()[:120], \"...\")"
]
},
{
@@ -169,10 +168,9 @@
"id": "12",
"metadata": {},
"source": [
- "### 1.3 TikZ-specific kwargs\n",
+ "### 1.4 Meshes, images and colorbars\n",
"\n",
- "The TikZ backend passes extra keyword arguments straight to `tikzfigure.draw()`. \n",
- "Use **`line_width=`** (not matplotlib's `linewidth=`) to control stroke thickness."
+ "Large colormap data would be slow and memory hungry as pgfplots coordinates, so meshes and images are drawn by Matplotlib and placed in the axis as images (`\\addplot graphics`). A colorbar is an axis of its own: an image of its colors, with pgfplots ticks and label. `files()` lists the images the figure refers to."
]
},
{
@@ -182,15 +180,19 @@
"metadata": {},
"outputs": [],
"source": [
- "canvas2, ax2 = Canvas.subplots(width=\"10cm\", ratio=0.5)\n",
- "ax2.plot(x, np.sin(x), color=\"navy\", line_width=0.5, label=\"thin\")\n",
- "ax2.plot(x, np.sin(x) + 0.5, color=\"steelblue\", line_width=1.5, label=\"medium\")\n",
- "ax2.plot(x, np.sin(x) + 1.0, color=\"royalblue\", line_width=3.0, label=\"thick\")\n",
- "ax2.set_xlabel(\"x\")\n",
- "ax2.set_title(\"Line width comparison\")\n",
+ "xx, yy = np.meshgrid(np.linspace(-2, 2, 120), np.linspace(-1, 1, 60))\n",
+ "field = np.exp(-(xx**2) - 4 * yy**2) * np.cos(4 * xx)\n",
+ "\n",
+ "canvas_mesh, ax_mesh = Canvas.subplots(width=\"10cm\", ratio=0.5)\n",
+ "ax_mesh.pcolormesh(xx, yy, field, cmap=\"RdBu_r\", vmin=-1, vmax=1)\n",
+ "ax_mesh.add_colorbar(label=r\"$\\phi$ [V]\") # of the mesh drawn last\n",
+ "ax_mesh.contour(xx, yy, field, levels=[-0.5, 0.5], colors=\"k\")\n",
+ "ax_mesh.set_xlabel(\"$x$\")\n",
+ "ax_mesh.set_ylabel(\"$y$\")\n",
"\n",
- "tikz2 = canvas2.render(backend=\"tikzfigure\")\n",
- "print(tikz2.generate_tikz())"
+ "tikz_mesh = canvas_mesh.render(backend=\"tikzfigure\")\n",
+ "print(list(tikz_mesh.files()))\n",
+ "tikz_mesh.show(transparent=False)"
]
},
{
@@ -198,9 +200,9 @@
"id": "14",
"metadata": {},
"source": [
- "### 1.4 Layer-aware TikZ output\n",
+ "### 1.5 Layer-aware TikZ output\n",
"\n",
- "Assign data to layers with `layer=N`. \n",
+ "Assign data to layers with `layer=N`.\n",
"The TikZ backend respects the layer filter — useful for generating incremental reveal figures (e.g. in Beamer)."
]
},
@@ -212,24 +214,23 @@
"outputs": [],
"source": [
"canvas3, ax3 = Canvas.subplots(width=\"10cm\", ratio=0.55)\n",
- "ax3.plot(x, np.sin(x), color=\"steelblue\", line_width=1.5, layer=0, label=\"sin\")\n",
- "ax3.plot(x, np.cos(x), color=\"tomato\", line_width=1.5, layer=1, label=\"cos\")\n",
+ "ax3.plot(x, np.sin(x), color=\"steelblue\", linewidth=1.5, layer=0, label=\"sin\")\n",
+ "ax3.plot(x, np.cos(x), color=\"tomato\", linewidth=1.5, layer=1, label=\"cos\")\n",
"ax3.plot(\n",
- " x, np.sin(x) * np.cos(x), color=\"seagreen\", line_width=1.0, layer=2, label=\"sin·cos\"\n",
+ " x, np.sin(x) * np.cos(x), color=\"seagreen\", linewidth=1.0, layer=2, label=\"sin·cos\"\n",
")\n",
"\n",
"# All layers available on the canvas\n",
"print(\"Available layers:\", canvas3.layers)\n",
"\n",
- "# Render only layer 0 — one \\draw command\n",
+ "# Render only layer 0 — one \\addplot\n",
"tikz_l0 = canvas3.render(backend=\"tikzfigure\", layers=[0])\n",
"print(\"\\n--- Layer 0 only ---\")\n",
- "print(f\"\\\\draw count: {tikz_l0.generate_tikz().count(chr(92) + 'draw')}\")\n",
+ "print(f\"\\\\addplot count: {tikz_l0.generate_tikz().count(chr(92) + 'addplot')}\")\n",
"\n",
- "# Render layers 0 and 1\n",
+ "# Layers 0 and 1\n",
"tikz_l01 = canvas3.render(backend=\"tikzfigure\", layers=[0, 1])\n",
- "print(\"\\n--- Layers 0 & 1 ---\")\n",
- "print(f\"\\\\draw count: {tikz_l01.generate_tikz().count(chr(92) + 'draw')}\")"
+ "print(f\"\\\\addplot count: {tikz_l01.generate_tikz().count(chr(92) + 'addplot')}\")"
]
},
{
@@ -237,9 +238,20 @@
"id": "16",
"metadata": {},
"source": [
- "### 1.5 Saving TikZ code to a file\n",
+ "### 1.6 Saving TikZ code to a file\n",
+ "\n",
+ "`savefig(\"figure.tikz\")` writes the `tikzpicture`, `savefig(\"figure.tex\")` a standalone document; both write the images the figure refers to next to the file. `savefig(\"figure.pdf\")` compiles it (requires `pdflatex`). In your document, load `pgfplots` and `\\input` the `.tikz` file:\n",
"\n",
- "You can embed the generated code directly in a LaTeX document:"
+ "```latex\n",
+ "\\usepackage{pgfplots}\n",
+ "\\pgfplotsset{compat=newest}\n",
+ "...\n",
+ "\\begin{figure}\n",
+ " \\centering\n",
+ " \\input{figure.tikz}\n",
+ " \\caption{My caption}\n",
+ "\\end{figure}\n",
+ "```"
]
},
{
@@ -249,21 +261,8 @@
"metadata": {},
"outputs": [],
"source": [
- "tikz_all = canvas3.render(backend=\"tikzfigure\")\n",
- "\n",
- "with open(\"figure.tex\", \"w\") as f:\n",
- " f.write(tikz_all.generate_tikz())\n",
- "\n",
- "print(\"Saved figure.tex\")\n",
- "\n",
- "# In your LaTeX document:\n",
- "# \\input{figure.tex}\n",
- "# or wrap it:\n",
- "# \\begin{figure}[h]\n",
- "# \\centering\n",
- "# \\input{figure.tex}\n",
- "# \\caption{My caption}\n",
- "# \\end{figure}"
+ "tikz_mesh.savefig(\"figure.tikz\")\n",
+ "print(\"Saved figure.tikz and\", \", \".join(tikz_mesh.files()))"
]
},
{
@@ -271,9 +270,9 @@
"id": "18",
"metadata": {},
"source": [
- "### 1.6 Rendering the figure (requires `pdflatex`)\n",
+ "### 1.7 Rendering the figure (requires `pdflatex`)\n",
"\n",
- "If `pdflatex` is installed, `tikz.show()` compiles the code and opens the PDF:"
+ "If `pdflatex` is installed, `tikz.show()` compiles the code and displays it:"
]
},
{
@@ -284,7 +283,7 @@
"outputs": [],
"source": [
"# Requires pdflatex:\n",
- "tikz_all.show(transparent=False)"
+ "tikz_grid.show(transparent=False)"
]
},
{
@@ -292,24 +291,65 @@
"id": "20",
"metadata": {},
"source": [
- "### 1.7 Canvas → TikZ limitations\n",
+ "### 1.8 Any Matplotlib figure\n",
+ "\n",
+ "The conversion works for every Matplotlib figure, not only canvases: `figure_to_tikz(fig)` converts a figure you drew with Matplotlib directly (it is what `render(backend=\"tikzfigure\")` uses), and `Canvas.from_matplotlib(fig)` imports one into a canvas to edit first."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "id": "21",
+ "metadata": {},
+ "outputs": [],
+ "source": [
+ "import matplotlib.pyplot as plt\n",
"\n",
- "| Feature | Supported? |\n",
+ "from maxplotlib.backends.tikzfigure import figure_to_tikz\n",
+ "\n",
+ "fig, (left, right) = plt.subplots(1, 2, figsize=(8, 3), layout=\"constrained\")\n",
+ "t = np.linspace(0, 10, 400)\n",
+ "left.semilogy(t, np.exp(0.4 * t), label=\"energy\")\n",
+ "left.semilogy(t, 0.5 * np.exp(0.4 * t), \"k--\", label=r\"fit, $\\gamma = 0.2$\")\n",
+ "left.axvspan(2, 6, color=\"gray\", alpha=0.2)\n",
+ "left.set_xlabel(\"$t$\")\n",
+ "left.legend()\n",
+ "right.scatter(\n",
+ " *np.random.default_rng(1).normal(size=(2, 200)), s=8, c=t[:200], cmap=\"viridis\"\n",
+ ")\n",
+ "right.set_aspect(\"equal\")\n",
+ "\n",
+ "tikz_mpl = figure_to_tikz(fig)\n",
+ "plt.close(fig)\n",
+ "tikz_mpl.show(transparent=False)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "id": "22",
+ "metadata": {},
+ "source": [
+ "### 1.9 What becomes what\n",
+ "\n",
+ "| Matplotlib | TikZ |\n",
"|---|---|\n",
- "| Line plots (`canvas.plot`) | ✅ |\n",
- "| Layer filtering | ✅ |\n",
- "| `line_width=` kwarg | ✅ |\n",
- "| Horizontal subplots (1×n) | ✅ |\n",
- "| `canvas.scatter`, `canvas.bar`, `canvas.barh` | ✅ |\n",
- "| `canvas.fill_between`, `canvas.errorbar` | ✅ |\n",
- "| Axis labels / titles | ✅ |\n",
+ "| Lines, markers, error bars, steps, stems | `\\addplot` with the same color, width, dashes and marks |\n",
+ "| Scatter plots | `\\addplot[only marks]`, one per color and size |\n",
+ "| Bars, fills, spans, polygons, patches | closed `\\addplot ... -- cycle` |\n",
+ "| Contour lines | `\\addplot` per level |\n",
+ "| Text, annotations, titles, labels | LaTeX text (mathtext `$...$` as math, the rest escaped) |\n",
+ "| Log axes, reversed axes, hidden spines, ticks | pgfplots axis options |\n",
+ "| Legends | `\\addlegendimage` + `\\addlegendentry`, in Matplotlib's order |\n",
+ "| Meshes, images, filled contours, quivers | images (`\\addplot graphics`) |\n",
+ "| Colorbars | an axis with the colors as an image |\n",
+ "| Polar or 3-D axes, symlog scales, figure legends | an image, with a `TikzConversionWarning` |\n",
"\n",
- "For unsupported primitives, the Canvas API raises `NotImplementedError`; use the direct `tikzfigure` API for advanced TikZ shapes (Part 2 below)."
+ "Scatter plots with more than `max_markers` points, collections of more than `max_items` styles and lines of more than `max_points` points (TeX memory is limited) are images too; `canvas.plot_tikzfigure(raster_dpi=600, max_markers=5000)` changes the limits and the image resolution."
]
},
{
"cell_type": "markdown",
- "id": "21",
+ "id": "23",
"metadata": {},
"source": [
"---\n",
@@ -327,7 +367,7 @@
},
{
"cell_type": "markdown",
- "id": "22",
+ "id": "24",
"metadata": {},
"source": [
"### 2.1 Drawing paths with `draw()`\n",
@@ -338,7 +378,7 @@
{
"cell_type": "code",
"execution_count": null,
- "id": "23",
+ "id": "25",
"metadata": {},
"outputs": [],
"source": [
@@ -356,7 +396,7 @@
},
{
"cell_type": "markdown",
- "id": "24",
+ "id": "26",
"metadata": {},
"source": [
"### 2.2 Straight line segments with `line()`\n",
@@ -368,7 +408,7 @@
{
"cell_type": "code",
"execution_count": null,
- "id": "25",
+ "id": "27",
"metadata": {},
"outputs": [],
"source": [
@@ -389,7 +429,7 @@
},
{
"cell_type": "markdown",
- "id": "26",
+ "id": "28",
"metadata": {},
"source": [
"### 2.3 Rectangles, circles, and arcs"
@@ -398,7 +438,7 @@
{
"cell_type": "code",
"execution_count": null,
- "id": "27",
+ "id": "29",
"metadata": {},
"outputs": [],
"source": [
@@ -431,7 +471,7 @@
},
{
"cell_type": "markdown",
- "id": "28",
+ "id": "30",
"metadata": {},
"source": [
"### 2.4 Nodes — text labels and markers\n",
@@ -442,7 +482,7 @@
{
"cell_type": "code",
"execution_count": null,
- "id": "29",
+ "id": "31",
"metadata": {},
"outputs": [],
"source": [
@@ -473,7 +513,7 @@
},
{
"cell_type": "markdown",
- "id": "30",
+ "id": "32",
"metadata": {},
"source": [
"### 2.5 Custom colours with `colorlet()`\n",
@@ -484,7 +524,7 @@
{
"cell_type": "code",
"execution_count": null,
- "id": "31",
+ "id": "33",
"metadata": {},
"outputs": [],
"source": [
@@ -508,7 +548,7 @@
},
{
"cell_type": "markdown",
- "id": "32",
+ "id": "34",
"metadata": {},
"source": [
"### 2.6 Filled paths and patterns\n",
@@ -519,7 +559,7 @@
{
"cell_type": "code",
"execution_count": null,
- "id": "33",
+ "id": "35",
"metadata": {},
"outputs": [],
"source": [
@@ -548,7 +588,7 @@
},
{
"cell_type": "markdown",
- "id": "34",
+ "id": "36",
"metadata": {},
"source": [
"### 2.7 Layers in `TikzFigure`\n",
@@ -560,7 +600,7 @@
{
"cell_type": "code",
"execution_count": null,
- "id": "35",
+ "id": "37",
"metadata": {},
"outputs": [],
"source": [
@@ -582,7 +622,7 @@
},
{
"cell_type": "markdown",
- "id": "36",
+ "id": "38",
"metadata": {},
"source": [
"### 2.8 Escaping to raw TikZ code\n",
@@ -593,7 +633,7 @@
{
"cell_type": "code",
"execution_count": null,
- "id": "37",
+ "id": "39",
"metadata": {},
"outputs": [],
"source": [
@@ -613,7 +653,7 @@
},
{
"cell_type": "markdown",
- "id": "38",
+ "id": "40",
"metadata": {},
"source": [
"### 2.9 Putting it all together — a complete figure\n",
@@ -624,7 +664,7 @@
{
"cell_type": "code",
"execution_count": null,
- "id": "39",
+ "id": "41",
"metadata": {},
"outputs": [],
"source": [
@@ -676,7 +716,7 @@
{
"cell_type": "code",
"execution_count": null,
- "id": "40",
+ "id": "42",
"metadata": {},
"outputs": [],
"source": [
@@ -686,7 +726,7 @@
},
{
"cell_type": "markdown",
- "id": "41",
+ "id": "43",
"metadata": {},
"source": [
"### 2.10 Embedding in a LaTeX document\n",
@@ -717,7 +757,7 @@
},
{
"cell_type": "markdown",
- "id": "42",
+ "id": "44",
"metadata": {},
"source": [
"---\n",
@@ -759,7 +799,7 @@
},
{
"cell_type": "markdown",
- "id": "43",
+ "id": "45",
"metadata": {},
"source": [
"## Part 1.8 — Canvas primitives supported by TikZ\n",
@@ -771,7 +811,7 @@
{
"cell_type": "code",
"execution_count": null,
- "id": "44",
+ "id": "46",
"metadata": {},
"outputs": [],
"source": [
@@ -798,7 +838,7 @@
{
"cell_type": "code",
"execution_count": null,
- "id": "45",
+ "id": "47",
"metadata": {},
"outputs": [],
"source": [
@@ -808,7 +848,7 @@
},
{
"cell_type": "markdown",
- "id": "46",
+ "id": "48",
"metadata": {},
"source": [
"The next cell shows the TikZ generated from the canvas. Unsupported primitives now raise\n",
@@ -817,7 +857,7 @@
},
{
"cell_type": "markdown",
- "id": "47",
+ "id": "49",
"metadata": {},
"source": [
"## Part 1.9 — More TikZ-supported primitives\n",
@@ -829,7 +869,7 @@
{
"cell_type": "code",
"execution_count": null,
- "id": "48",
+ "id": "50",
"metadata": {},
"outputs": [],
"source": [
@@ -849,7 +889,7 @@
{
"cell_type": "code",
"execution_count": null,
- "id": "49",
+ "id": "51",
"metadata": {},
"outputs": [],
"source": [
diff --git a/tutorials/tutorial_15_tikzfigure_subplots.ipynb b/tutorials/tutorial_15_tikzfigure_subplots.ipynb
index 16daecd..c206c8d 100644
--- a/tutorials/tutorial_15_tikzfigure_subplots.ipynb
+++ b/tutorials/tutorial_15_tikzfigure_subplots.ipynb
@@ -7,7 +7,7 @@
"source": [
"# Tutorial 15 - TikzFigure Subplots Tutorial\n",
"\n",
- "This tutorial demonstrates how to create side-by-side subplots using the `tikzfigure` backend.\n",
+ "This tutorial demonstrates how to create subplots using the `tikzfigure` backend: side by side, stacked, or in a grid.\n",
"\n",
"## Basic 1×2 Layout\n",
"\n",
@@ -64,8 +64,7 @@
"source": [
"## Inspecting the generated subplot dimensions\n",
"\n",
- "The exported TikZ splits the available width across the columns and keeps the requested overall height.\n",
- "Printing the `\\\\nextgroupplot[...]` lines makes that explicit."
+ "Every subplot is a pgfplots `axis` with the position and size Matplotlib gives it: `at=` is its lower left corner in the figure, `width=` and `height=` its box (`scale only axis`). Printing the `\\begin{axis}[...]` lines makes that explicit."
]
},
{
@@ -79,10 +78,17 @@
"subplot_code = tikz.generate_tikz()\n",
"\n",
"for line in subplot_code.splitlines():\n",
- " if \"nextgroupplot\" in line:\n",
- " print(line.strip())\n",
- "\n",
- "# With width=\"10cm\" and ratio=0.3, each subplot gets its own width entry."
+ " if \"\\\\begin{axis}\" in line:\n",
+ " options = line.strip()\n",
+ " print(\n",
+ " [\n",
+ " part\n",
+ " for part in options.split(\", \")\n",
+ " if part.startswith((\"at=\", \"width=\", \"height=\"))\n",
+ " ]\n",
+ " )\n",
+ "\n",
+ "# With width=\"10cm\" and ratio=0.3, the two axes share the figure width."
]
},
{
@@ -130,13 +136,38 @@
"cell_type": "markdown",
"id": "7",
"metadata": {},
+ "source": [
+ "## Grids and stacked layouts\n",
+ "\n",
+ "Any `nrows × ncols` layout converts, with shared axes and hidden tick labels as in Matplotlib."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "id": "8",
+ "metadata": {},
+ "outputs": [],
+ "source": [
+ "canvas, axes = Canvas.subplots(nrows=2, ncols=2, width=\"12cm\", ratio=0.6)\n",
+ "x = np.linspace(0, 2 * np.pi, 100)\n",
+ "for index, ax in enumerate(np.ravel(axes)):\n",
+ " ax.plot(x, np.sin((index + 1) * x), color=f\"C{index}\")\n",
+ " ax.set_title(f\"sin({index + 1}x)\")\n",
+ "\n",
+ "canvas.show(backend=\"tikzfigure\")"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "id": "9",
+ "metadata": {},
"source": [
"## Important Notes\n",
"\n",
- "- Only **horizontal layouts (1×n)** are supported with tikzfigure backend\n",
- "- Vertical/grid layouts (nrows > 1) will raise an error\n",
- "- Use the direct tikzfigure API for complex layouts or grids\n",
- "- Each subplot's title becomes a pgfplots `title=` entry in the generated LaTeX output"
+ "- Any layout converts: rows, columns and grids, with shared and twin axes, as Matplotlib lays them out\n",
+ "- Each subplot's title becomes a pgfplots `title=` entry in the generated LaTeX output\n",
+ "- Meshes, images and colorbars are included as images; see Tutorial 07 for what becomes what"
]
}
],