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") ![](README_files/figure-commonmark/cell-14-output-1.png) -### Horizontal Subplots with TikZ Backend +### Subplots and Meshes with the TikZ Backend -The tikzfigure backend supports creating side-by-side subplots (1×n -layouts): +The tikzfigure backend draws the canvas with Matplotlib and converts the +drawn figure into pgfplots axes, so every layout converts (rows, columns, +grids, twin axes), with LaTeX text, legends and colorbars. Lines, +markers, bars and text are pgfplots code; meshes and images are included +as images: ``` python x = np.linspace(0, 2 * np.pi, 200) -canvas, (ax1, ax2) = Canvas.subplots(ncols=2, width="10cm", ratio=0.3) +canvas, (ax1, ax2) = Canvas.subplots(ncols=2, width="12cm", ratio=0.45) -ax1.plot(x, np.sin(x), color="royalblue") -ax1.set_title("sin(x)") +ax1.plot(x, np.sin(x), color="royalblue", label="$\\sin x$") +ax1.plot(x, np.cos(x), color="tomato", label="$\\cos x$") +ax1.set_title("Lines") +ax1.set_legend(True) -ax2.plot(x, np.cos(x), color="tomato") -ax2.set_title("cos(x)") +xx, yy = np.meshgrid(x, x) +ax2.pcolormesh(xx, yy, np.sin(xx) * np.cos(yy), cmap="RdBu_r") +ax2.add_colorbar(label="$\\sin x \\cos y$") +ax2.set_title("A mesh") -canvas.suptitle("Trigonometric Functions") -canvas.show(backend="tikzfigure") # Generates LaTeX subfigures +canvas.show(backend="tikzfigure") # compiles with pdflatex ```
@@ -248,9 +254,11 @@ Figure 2
-**Note:** Only horizontal layouts (1×n) are currently supported with the -tikzfigure backend. Vertical/grid layouts will raise -`NotImplementedError`. See the tutorials for more examples. +`canvas.render(backend="tikzfigure").savefig("figure.tikz")` writes the +code for `\\input` in a LaTeX document, with the images next to it. Any +Matplotlib figure converts the same way with +`maxplotlib.backends.tikzfigure.figure_to_tikz(fig)`. See the tutorials +for more examples. ### Terminal Backend with plotext 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" ] } ],