diff --git a/README.md b/README.md index aab834e..9847d3a 100644 --- a/README.md +++ b/README.md @@ -2,124 +2,26 @@ Live training curves in Jupyter, Colab and the VS Code / Cursor interactive window, with a tqdm bar underneath, at (almost) no cost to the training loop. -[![Open in Colab](https://colab.research.google.com/assets/colab-badge.svg)](https://colab.research.google.com/github/ARENA-education/liveplot/blob/demo/examples/demo.ipynb) [![tests](https://github.com/ARENA-education/liveplot/actions/workflows/tests.yml/badge.svg)](https://github.com/ARENA-education/liveplot/actions/workflows/tests.yml) - -Documentation: **[arena-education.github.io/liveplot](https://arena-education.github.io/liveplot/)** (this README, the API, and the examples; built from the `.md` files on every push). - -Try it in Colab with the badge above: that notebook is [`examples/demo.py`](examples/demo.py), a cell-by-cell tour of the features, which CI converts with jupytext and publishes to the `demo` branch on every push to `main`. +[![docs](https://img.shields.io/badge/docs-arena--education.github.io%2Fliveplot-blue)](https://arena-education.github.io/liveplot/) [![Open in Colab](https://colab.research.google.com/assets/colab-badge.svg)](https://colab.research.google.com/github/ARENA-education/liveplot/blob/demo/examples/demo.ipynb) [![tests](https://github.com/ARENA-education/liveplot/actions/workflows/tests.yml/badge.svg)](https://github.com/ARENA-education/liveplot/actions/workflows/tests.yml) ![training loss every step, eval loss and accuracy every 50 steps, on one panel with two y-axes](docs/demo.gif) -The GIF is cell 2 of [`examples/demo.py`](examples/demo.py), a tour of the features in a cell-separated file: open it in the VS Code / Cursor interactive window and run it cell by cell. That cell trains a tiny numpy classifier, logging the training loss every step and the eval loss and accuracy every 50 steps, laid out as `"loss | eval_loss eval_acc"`. The GIF itself was recorded by the library: `examples/make_gif.py` runs the same loop as a plain script with `record="docs/demo.gif"`. - -```python -from liveplot import LivePlot - -plot = LivePlot(range(num_steps)) # a tqdm bar + a live plot -for step in plot: - loss, acc = train_step() - plot.log(loss=loss, acc=acc) # step is implicit; the bar's postfix shows the latest values -``` - -Metrics are discovered from what you log. With no layout given they all share one panel, with a legend. To split them up, give panel strings: - -```python -plot = LivePlot(range(num_steps), "loss", "return | entropy", "lossD lossG | acc") -``` - -Each string is one panel. Names separated by spaces share the left y-axis; names after a `|` go on a right-hand y-axis. Anything you log that no string mentions gets a panel of its own. Every panel has a legend. - -For nested loops, create the plot once and wrap the inner loop with `plot(...)`; the count, and so the x-axis, continues across epochs and you get one tqdm bar per epoch, as with tqdm: - -```python -plot = LivePlot("loss", "acc", total=epochs * len(loader)) -for epoch in range(epochs): - for imgs, labels in plot(loader, desc=f"epoch {epoch}"): # kwargs go to tqdm - plot.log(loss=train_step(imgs, labels)) - plot.log(acc=evaluate()) # metrics can have different cadences -plot.finish() # or wrap the whole thing in `with LivePlot(...) as plot:` -``` - -If you'd rather supply the x values yourself, use the context manager and pass the step to `log`: - -```python -with LivePlot("loss", "acc", total=num_steps) as plot: - for step in range(num_steps): - plot.log(step, loss=train_step()) -``` - -Already have a tqdm bar? Pass it as the iterable and it is reused instead of wrapped: `LivePlot(tqdm(loader), "loss")`. - -## Metrics logged at different rates - -Nothing special is needed. Every metric keeps its own list of `(x, value)` points, and each `log` call stamps its metrics with the current x. A metric logged every step gets a point per step; one logged every 50 steps, or once per epoch after the inner loop, gets a point wherever the count was at the time and is drawn as a line through those points. So in `"loss | eval_loss eval_acc"` the left axis fills in continuously while the right axis grows a point every 50 steps, all against the same x. - -## The x-axis - -It follows tqdm: the plot counts items consumed and never looks at their values. - - x = initial + n * unit_scale - -`n` is 0 inside the loop body for the first item, like `for step in range(N)`. `initial` shifts the start (`initial=1` for 1-based, `initial=10` to plot `range(10, 20)` at its values). `unit` and `unit_scale` relabel and rescale it: `unit="examples", unit_scale=batch_size` plots against examples seen, and the tqdm bar shows the same numbers. `total`, in items like tqdm's, fixes the x range; the single-loop form takes it from `len(iterable)`. An explicit `plot.log(step, ...)` uses that x for that call only, like wandb's `step=`. - -## Install - ``` pip install git+https://github.com/ARENA-education/liveplot.git ``` -Only `matplotlib` is required. `ipython` is needed for the live display and `tqdm` for the bar; any notebook has both, and without them the plot silently just collects `plot.data`. - -## How it works - -The training thread only appends numbers (about 40 µs per `log`). A separate render process owns the matplotlib figure, redraws it at most once per `refresh_seconds` (default 1; `0` means on every arrival), and sends back PNG bytes that get swapped into a fixed output cell. The output is a plain image, so it behaves identically in Jupyter, Colab, VS Code and Cursor: no widgets, no JavaScript, no CDN. - -Interrupting the cell is safe. Jupyter sends its interrupt to every process the kernel started; the render process ignores it, so you get a frozen plot with `plot.data` intact. A plot that is dropped without `finish()` shuts its process down when garbage collected, and the process exits by itself if the notebook kernel dies. - -## Options - -| | | -|---|---| -| `total`, `initial`, `unit`, `unit_scale` | tqdm's arguments, with tqdm's meaning; they define the x-axis (see above) | -| `refresh_seconds` | minimum time between redraws (default 1.0). Points arriving in between are batched into the next frame. `0` redraws whenever new data arrives, as fast as rendering allows (roughly 0.15 s per frame at a few thousand points), and costs nothing while idle. | -| `max_cols`, `rows`, `cols` | grid shape; `max_cols=None` gives a near-square grid | -| `progress`, `desc` | disable the bundled tqdm bars, or give the single-loop form's bar a description | -| `cell_size`, `dpi` | size of each panel in inches, and PNG resolution | -| `record` | `True` keeps every rendered frame in `plot.frames`; a path such as `"run.gif"` also writes an animated GIF at `finish()`. `plot.save_gif(path, speedup=1.0)` does it on demand. Works outside a notebook too. | - -For labels or fixed ranges, use a dict instead of a string for that panel: - -```python -{"metrics": ["acc"], "ylim": (0, 1), "ylabel": "test accuracy", "xlabel": "epoch"} -``` - -Allowed keys: `title`, `metrics`, `secondary`, `xlabel`, `ylabel`, `ylabel2`, `xlim`, `ylim`, `ylim2`, `axhlines`, `axhlines2`, `smooth`, `yscale`, `yscale2`. Where a key has a matplotlib counterpart it uses matplotlib's name (`ax.set(title=..., xlabel=..., xlim=..., yscale=...)`); the `2` suffix means the right-hand axis. - -## Reference lines - -The names follow matplotlib. A dashed horizontal line with a legend entry, for the level a curve should reach or beat: - ```python -LivePlot(range(N), {"metrics": ["loss"], "axhlines": {"uniform": math.log(d_vocab), "unigram": 7.35}}) -plot.axhline(500, "solved", metric="return") # at run time; goes on the axis of that metric's panel -plot.axvline(label="lr drop") # a dotted vertical line on every panel at the current x -plot.axhline(0.9, "target", metric="acc", color="red", linestyle="-") # extra kwargs go to the matplotlib artist -``` - -In a panel dict, `axhlines` takes `{label: y}`, a list of values, or a list of `axhline` kwargs such as `dict(y=0.9, label="target", color="red")`; `axhlines2` is the same for the right-hand axis. - -## Smoothing and log axes - -Per-step losses are noisy. `smooth=0.9` on a panel draws each of its curves through wandb's default smoothing, the time-weighted exponential moving average, with the same 0 to 1 weight as wandb's smoothing slider and the raw values faded behind. `LivePlot(..., smooth=0.9)` makes that the default for every panel, and `"smooth": 0` on a panel opts out. `yscale="log"` (and `yscale2` for the right axis) gives a log axis. +from liveplot import LivePlot -```python -LivePlot(range(N), {"metrics": ["loss"], "smooth": 0.9, "yscale": "log"}, "acc") +plot = LivePlot(loader, "loss | acc", "lr") # wrap the loop like tqdm; "|" puts acc on a right-hand axis +plot["acc"].set_ylim(0, 1) # configure like matplotlib +for batch in plot: + loss, acc = train_step(batch) + plot.log(loss=loss, acc=acc, lr=lr) # log like wandb ``` -`plot.figure()` returns a matplotlib Figure of the plot as it stands, built independently of the live renderer, for `fig.savefig(...)`, a title, or any other tweak. - -`plot.log` accepts keywords, an explicit step (`plot.log(step, loss=...)`), or a dict (`plot.log(step, {"loss": ...})`). Values can be anything `float()` accepts, including one-element tensors. `plot.data` holds the full history as `{metric: (steps, values)}` and `plot.latest` the most recent value of each. +It wraps iterables like tqdm, logs like wandb and is configured like matplotlib, so the names are ones you already know. Rendering happens in a separate process and the output is a plain image, so it behaves the same everywhere with no widgets or JavaScript, and interrupting a cell is safe. -## Credits +**[Documentation](https://arena-education.github.io/liveplot/)**: the [guide](https://arena-education.github.io/liveplot/guide/) (nested loops, the x-axis, reference lines, smoothing, recording GIFs), the [API](https://arena-education.github.io/liveplot/api/), and the [examples](https://arena-education.github.io/liveplot/examples/). The tour in [`examples/demo.py`](examples/demo.py) runs cell by cell in the interactive window or [in Colab](https://colab.research.google.com/github/ARENA-education/liveplot/blob/demo/examples/demo.ipynb). -The per-panel label/limit options and the grid-shape rule are adapted from Tyler Lum's [live_plotter](https://github.com/tylerlum/live_plotter) (MIT); see `THIRD_PARTY_LICENSES.md`. +MIT. Per-panel option names and the grid rule are adapted from Tyler Lum's [live_plotter](https://github.com/tylerlum/live_plotter); see `THIRD_PARTY_LICENSES.md`. diff --git a/docs/axes_api.gif b/docs/axes_api.gif new file mode 100644 index 0000000..914844c Binary files /dev/null and b/docs/axes_api.gif differ diff --git a/docs/guide.md b/docs/guide.md new file mode 100644 index 0000000..6190ab9 --- /dev/null +++ b/docs/guide.md @@ -0,0 +1,117 @@ +# Guide + +The one-liner first, then everything it can grow into. + +```python +from liveplot import LivePlot + +plot = LivePlot(range(num_steps)) # a tqdm bar + a live plot +for step in plot: + loss, acc = train_step() + plot.log(loss=loss, acc=acc) # step is implicit; the bar's postfix shows the latest values +``` + +Metrics are discovered from what you log. With no layout given they all share one panel, with a legend. To split them up, give panel strings: + +```python +plot = LivePlot(range(num_steps), "loss", "return | entropy", "lossD lossG | acc") +``` + +Each string is one panel. Names separated by spaces share the left y-axis; names after a `|` go on a right-hand y-axis. Anything you log that no string mentions gets a panel of its own. Every panel has a legend. + +For nested loops, create the plot once and wrap the inner loop with `plot(...)`; the count, and so the x-axis, continues across epochs and you get one tqdm bar per epoch, as with tqdm: + +```python +plot = LivePlot("loss", "acc", total=epochs * len(loader)) +for epoch in range(epochs): + for imgs, labels in plot(loader, desc=f"epoch {epoch}"): # kwargs go to tqdm + plot.log(loss=train_step(imgs, labels)) + plot.log(acc=evaluate()) # metrics can have different cadences +plot.finish() # or wrap the whole thing in `with LivePlot(...) as plot:` +``` + +If you'd rather supply the x values yourself, use the context manager and pass the step to `log`: + +```python +with LivePlot("loss", "acc", total=num_steps) as plot: + for step in range(num_steps): + plot.log(step, loss=train_step()) +``` + +Already have a tqdm bar? Pass it as the iterable and it is reused instead of wrapped: `LivePlot(tqdm(loader), "loss")`. + +## Metrics logged at different rates + +Nothing special is needed. Every metric keeps its own list of `(x, value)` points, and each `log` call stamps its metrics with the current x. A metric logged every step gets a point per step; one logged every 50 steps, or once per epoch after the inner loop, gets a point wherever the count was at the time and is drawn as a line through those points. So in `"loss | eval_loss eval_acc"` the left axis fills in continuously while the right axis grows a point every 50 steps, all against the same x. + +## The x-axis + +It follows tqdm: the plot counts items consumed and never looks at their values. + + x = initial + n * unit_scale + +`n` is 0 inside the loop body for the first item, like `for step in range(N)`. `initial` shifts the start (`initial=1` for 1-based, `initial=10` to plot `range(10, 20)` at its values). `unit` and `unit_scale` relabel and rescale it: `unit="examples", unit_scale=batch_size` plots against examples seen, and the tqdm bar shows the same numbers. `total`, in items like tqdm's, fixes the x range; the single-loop form takes it from `len(iterable)`. An explicit `plot.log(step, ...)` uses that x for that call only, like wandb's `step=`. + +## How it works + +The training thread only appends numbers (about 40 µs per `log`). A separate render process owns the matplotlib figure, redraws it at most once per `refresh_seconds` (default 1; `0` means on every arrival), and sends back PNG bytes that get swapped into a fixed output cell. The output is a plain image, so it behaves identically in Jupyter, Colab, VS Code and Cursor: no widgets, no JavaScript, no CDN. + +Interrupting the cell is safe. Jupyter sends its interrupt to every process the kernel started; the render process ignores it, so you get a frozen plot with `plot.data` intact. A plot that is dropped without `finish()` shuts its process down when garbage collected, and the process exits by itself if the notebook kernel dies. + +## Options + +| | | +|---|---| +| `total`, `initial`, `unit`, `unit_scale` | tqdm's arguments, with tqdm's meaning; they define the x-axis (see above) | +| `refresh_seconds` | minimum time between redraws (default 1.0). Points arriving in between are batched into the next frame. `0` redraws whenever new data arrives, as fast as rendering allows (roughly 0.15 s per frame at a few thousand points), and costs nothing while idle. | +| `max_cols`, `rows`, `cols` | grid shape; `max_cols=None` gives a near-square grid | +| `progress`, `desc` | disable the bundled tqdm bars, or give the single-loop form's bar a description | +| `cell_size`, `dpi` | size of each panel in inches, and PNG resolution | +| `record` | `True` keeps every rendered frame in `plot.frames`; a path such as `"run.gif"` also writes an animated GIF at `finish()`. `plot.save_gif(path, speedup=1.0)` does it on demand. Works outside a notebook too. | + +## Titles, labels, limits: matplotlib's names + +Panels and axes are addressed and configured with the setters you already know from matplotlib, before or during the loop: + +```python +plot = LivePlot(loader, "loss | acc", "lr") + +plot["acc"].set_ylim(0, 1) # plot[metric] is the y-axis holding that metric (left or right) +plot["acc"].set_ylabel("test accuracy") +plot["loss"].axhline(0.1, label="target", color="red") # extra kwargs go to the matplotlib artist + +plot.panels[0].set_title("training") # a panel: set_title, set_xlabel, set_xlim, set_smooth, axvline +plot.panels[0].set_xlim(0, 10_000) +plot.panels[0].axvline(label="lr drop") # at the current x + +plot.set(xlabel="examples", smooth=0.9) # on the plot itself: every panel, like Axes.set(**kwargs) +``` + +`plot[metric]` finds the right-hand axis too, so there is no separate spelling for it; `plot.panels[i].left` / `.right` name the two axes explicitly. Everything on an axis takes matplotlib's name and signature: `set_ylabel`, `set_ylim` (two arguments or a tuple), `set_yscale("log")`, `axhline`, `set(**kwargs)`, and the panel's own setters are reachable from it as they would be on an `Axes`. A setter called mid-run just triggers a re-layout on the next frame. + +The dict form, `{"metrics": ["acc"], "ylim": (0, 1), ...}`, is still accepted in the layout list but the setters are the intended way. + +`plot.figure()` returns a matplotlib Figure of the plot as it stands, built independently of the live renderer, for `fig.savefig(...)`, a title, or any other tweak. + +`plot.log` accepts keywords, an explicit step (`plot.log(step, loss=...)`), or a dict (`plot.log(step, {"loss": ...})`). Values can be anything `float()` accepts, including one-element tensors. `plot.data` holds the full history as `{metric: (steps, values)}` and `plot.latest` the most recent value of each. + +## Reference lines + +Matplotlib's `axhline` / `axvline`, on an axis, a panel, or the whole plot. Horizontal lines get a legend entry; extra kwargs go to the artist. + +```python +plot["loss"].axhline(math.log(d_vocab), "uniform") # the level a curve should beat +plot["return"].axhline(500, "solved", color="green") +plot.axvline(label="lr drop") # every panel, at the current x +plot.panels[1].axvline(2000, "checkpoint", linestyle="-") # one panel, at a given x +``` + +## Smoothing and log axes + +Per-step losses are noisy. `set_smooth(0.9)` draws each curve of a panel through wandb's default smoothing, the [time-weighted exponential moving average](https://docs.wandb.ai/models/app/features/panels/line-plot/smoothing), with the same 0 to 1 weight as wandb's smoothing slider and the raw values faded behind. On the plot it applies to every panel; `set_smooth(0)` turns it off. `set_yscale("log")` on an axis, a panel, or the plot gives log axes. + +```python +plot.set_smooth(0.9) +plot["loss"].set_yscale("log") +``` + diff --git a/examples/demo.py b/examples/demo.py index 782f541..0b4f9e1 100644 --- a/examples/demo.py +++ b/examples/demo.py @@ -101,11 +101,15 @@ def train_classifier(steps=600, batch_size=64, lr=0.5, eval_every=50, seed=0, pa # 3. Nested loops (epochs x batches). Create the plot once and wrap the INNER loop with plot(...): # the count, and so the x-axis, continues across epochs and you get one tqdm bar per epoch, # exactly as with tqdm. `total` is in items, like tqdm's, and fixes the x range up front. -# A dict instead of a string pins the accuracy axis to [0, 1] and labels it. +# Titles, labels and limits use matplotlib's own setter names, addressed by metric or by panel. if MAIN: epochs, loader = 3, [None] * 60 - plot = LivePlot("loss", {"metrics": ["acc"], "ylim": (0, 1), "ylabel": "test accuracy"}, total=epochs * len(loader)) + plot = LivePlot("loss", "acc", total=epochs * len(loader)) + plot["acc"].set_ylim(0, 1) # plot[metric] is the y-axis holding that metric + plot["acc"].set_ylabel("test accuracy") + plot.panels[0].set_title("training loss") + plot.set(xlabel="batches") # every panel, like Axes.set(**kwargs) for epoch in range(epochs): for batch in plot(loader, desc=f"epoch {epoch}"): plot.log(loss=math.exp(-plot.n / 60) + 0.05 * random.random()) @@ -126,6 +130,9 @@ def train_classifier(steps=600, batch_size=64, lr=0.5, eval_every=50, seed=0, pa plot.log(step, acc=min(1.0, step / 300)) if step >= 200: plot.log(step, lr=1e-3 * (400 - step) / 200) + if step == 200: + plot.axvline(label="lr decay starts") # a vertical line on every panel, at the current x + plot["acc"].axhline(0.9, "target") # a horizontal reference line on acc's axis slow() # %% diff --git a/liveplot/liveplot.py b/liveplot/liveplot.py index a3b3859..7a687a6 100644 --- a/liveplot/liveplot.py +++ b/liveplot/liveplot.py @@ -46,17 +46,23 @@ {"metrics": ["acc"], "ylim": (0, 1), "ylabel": "test accuracy", "xlabel": "epoch"} -(allowed keys: title, metrics, secondary, xlabel, ylabel, ylabel2, xlim, ylim, ylim2, -axhlines, axhlines2, smooth, yscale, yscale2; the names follow matplotlib's -`Axes.set(...)` keywords, with a `2` suffix for the right-hand axis). Reference -lines follow matplotlib too: `axhlines={"uniform": 10.8}` in a panel (or a list of -`axhline` kwargs) draws a dashed line at that level with a legend entry, -`plot.axhline(y, label, metric=...)` does the same at run time, and -`plot.axvline(label="lr drop")` draws a dotted vertical line on every panel at the -current x. Extra kwargs go to the artists. Smoothing: `smooth=0.9` on a panel (or -on LivePlot, as the default for every panel) draws each curve as wandb's -time-weighted EMA with that weight, with the raw values faded behind it. Log axes: -`yscale="log"` (`yscale2` for the right axis). +(the dict form; the setters below are the nicer way). Panels and axes are addressed +and configured with matplotlib's own names, before or during the loop: + + plot["acc"].set_ylim(0, 1) # the y-axis holding a metric (left or right) + plot["acc"].set_ylabel("test accuracy") + plot["loss"].axhline(0.1, label="target", color="red") + plot.panels[0].set_title("training") # a panel: title, xlabel, xlim, axvline, smooth + plot.panels[0].right.set_yscale("log") # .left / .right are the panel's y-axes + plot.set(xlabel="examples", smooth=0.9) # on the plot: every panel, like Axes.set + +Available on axes: set_ylabel, set_ylim, set_yscale, axhline, set(**kw); on panels: +set_title, set_xlabel, set_xlim, set_smooth, axvline, set(**kw), plus the left +axis's setters; on the plot: all of these for every panel, `axhline` (every left +axis, or `metric=` for one) and `axvline` (every panel). Extra kwargs on +`axhline` / `axvline` go to the matplotlib artist. Smoothing is wandb's +time-weighted EMA with the same 0 to 1 weight (`smooth=0.9`), raw values faded +behind. How it works: the training thread only appends numbers (~40 us per `log`). A separate *render process* owns the matplotlib figure, redraws it at most once per @@ -106,7 +112,7 @@ import weakref _PANEL_KEYS = {"title", "metrics", "secondary", "xlabel", "ylabel", "ylabel2", "xlim", "ylim", "ylim2", "axhlines", "axhlines2", - "smooth", "yscale", "yscale2"} + "axvlines", "smooth", "yscale", "yscale2"} _REF_LINE_STYLE = {"linestyle": "--", "linewidth": 1, "color": "0.45"} # defaults for axhline / axvline artists @@ -163,6 +169,7 @@ def _normalise_panel(spec) -> dict: "ylim2": spec.get("ylim2"), "axhlines": _normalise_axhlines(spec.get("axhlines")), # reference lines on the left axis (see _normalise_axhlines) "axhlines2": _normalise_axhlines(spec.get("axhlines2")), # ... and on the right axis + "axvlines": [dict(kw) for kw in spec.get("axvlines") or []], # vertical lines on this panel only (axvline kwargs) "smooth": spec.get("smooth"), # TWEMA weight in [0, 1); None = plot-wide default; 0 = off "yscale": spec.get("yscale", "linear"), # "linear" or "log", left axis "yscale2": spec.get("yscale2", "linear"), # ... right axis @@ -293,6 +300,9 @@ def set_layout(self, panels): for ax in axes[n:]: ax.set_visible(False) self.axes = list(axes[:n]) + for ax, panel in zip(self.axes, panels): + for kw in panel["axvlines"]: + self._draw_axvline(kw, [ax]) for kw in self.axvlines: self._draw_axvline(kw) self.fig.tight_layout() @@ -301,10 +311,10 @@ def add_axvline(self, kw: dict): self.axvlines.append(kw) self._draw_axvline(kw) - def _draw_axvline(self, kw): + def _draw_axvline(self, kw, axes=None): label = kw.get("label") style = {**_REF_LINE_STYLE, "linestyle": ":", **{k: v for k, v in kw.items() if k != "label"}} - for ax in self.axes: + for ax in self.axes if axes is None else axes: ax.axvline(**style) if label: ax.text(kw["x"], 0.98, f" {label}", transform=ax.get_xaxis_transform(), va="top", ha="left", fontsize=7, color="0.3", rotation=90) @@ -434,6 +444,149 @@ def _collect_loop(plot_ref, outbox, proc, done): done.set() +# --------------------------------------------------------------------------- panels and axes, addressed like matplotlib + + +def _lim(a, b): + """Accept matplotlib's forms: set_xlim((lo, hi)), set_xlim(lo, hi), set_xlim(lo) / set_xlim(right=hi) not supported.""" + if b is None and isinstance(a, (tuple, list)): + a, b = a + assert a is not None and b is not None, "give both limits, e.g. set_ylim(0, 1)" + return (float(a), float(b)) + + +class _Axis: + """ + One y-axis of a panel (left, or the right-hand `secondary` one), with matplotlib's `Axes` names: + set_ylabel / set_ylim / set_yscale / axhline / set(**kwargs). Panel-level setters (title, xlabel, + xlim, axvline, smooth) are available here too, as they are on a matplotlib Axes. + """ + + def __init__(self, plot, index: int, right: bool): + self._plot, self._index, self._right = plot, index, right + + @property + def panel(self): + return Panel(self._plot, self._index) + + @property + def metrics(self) -> list: + return list(self._plot._specs[self._index]["secondary" if self._right else "metrics"]) + + def _set(self, key, value): + self._plot._specs[self._index][key + ("2" if self._right else "")] = value + self._plot._send_layout() + + def set_ylabel(self, ylabel): + self._set("ylabel", str(ylabel)) + + def set_ylim(self, bottom=None, top=None): + self._set("ylim", _lim(bottom, top)) + + def set_yscale(self, value): + assert value in ("linear", "log"), f"yscale must be 'linear' or 'log', got {value!r}" + self._set("yscale", value) + + def axhline(self, y, label=None, **kwargs): + """Like `Axes.axhline`: a horizontal reference line on this axis, with `label` in the legend.""" + line = {**kwargs, "y": float(y), "label": str(label) if label is not None else f"{float(y):g}"} + self._plot._specs[self._index]["axhlines2" if self._right else "axhlines"].append(line) + self._plot._send_layout() + + def set(self, **kwargs): + """Like `Axes.set`: `ax.set(ylabel="loss", ylim=(0, 1), yscale="log", title=...)`.""" + for key, value in kwargs.items(): + getattr(self, f"set_{key}")(value) + + # panel-level setters, so plot["acc"].set_title(...) works as it would on a matplotlib Axes + def set_title(self, label): + self.panel.set_title(label) + + def set_xlabel(self, xlabel): + self.panel.set_xlabel(xlabel) + + def set_xlim(self, left=None, right=None): + self.panel.set_xlim(left, right) + + def axvline(self, x=None, label=None, **kwargs): + self.panel.axvline(x, label, **kwargs) + + def set_smooth(self, weight): + self.panel.set_smooth(weight) + + +class Panel: + """ + One panel of the grid: `plot.panels[i]`, or `plot[metric].panel`. `.left` / `.right` are its y-axes + (`.right` exists only if the panel has secondary metrics). y-setters here act on the left axis. + """ + + def __init__(self, plot, index: int): + self._plot, self._index = plot, index + + @property + def spec(self) -> dict: + return self._plot._specs[self._index] + + @property + def metrics(self) -> list: + return list(self.spec["metrics"]) + list(self.spec["secondary"]) + + @property + def left(self) -> _Axis: + return _Axis(self._plot, self._index, right=False) + + @property + def right(self) -> _Axis: + assert self.spec["secondary"], "this panel has no right-hand axis (no metrics after '|')" + return _Axis(self._plot, self._index, right=True) + + def _set(self, key, value): + self.spec[key] = value + self._plot._send_layout() + + def set_title(self, label): + self._set("title", str(label)) + + def set_xlabel(self, xlabel): + self._set("xlabel", str(xlabel)) + + def set_xlim(self, left=None, right=None): + self._set("xlim", _lim(left, right)) + + def set_smooth(self, weight): + """wandb-style smoothing weight in [0, 1) for every curve on this panel; 0 turns it off.""" + assert 0 <= weight < 1, f"smooth must be a weight in [0, 1), got {weight!r}" + self._set("smooth", weight) + + def axvline(self, x=None, label=None, **kwargs): + """Like `Axes.axvline`, on this panel only: a vertical line at `x` (default: the current step).""" + line = {**kwargs, "x": float(self._plot.step if x is None else x), "label": label} + self.spec["axvlines"].append(line) + self._plot._send_layout() + + def set(self, **kwargs): + """Like `Axes.set`: `panel.set(title="training", xlabel="examples", ylim=(0, 1))`.""" + for key, value in kwargs.items(): + getattr(self, f"set_{key}")(value) + + # y-setters act on the left axis, as on a matplotlib Axes + def set_ylabel(self, ylabel): + self.left.set_ylabel(ylabel) + + def set_ylim(self, bottom=None, top=None): + self.left.set_ylim(bottom, top) + + def set_yscale(self, value): + self.left.set_yscale(value) + + def axhline(self, y, label=None, **kwargs): + self.left.axhline(y, label, **kwargs) + + def __repr__(self): + return f"Panel({self._index}: {' '.join(self.spec['metrics'])}{' | ' + ' '.join(self.spec['secondary']) if self.spec['secondary'] else ''})" + + # --------------------------------------------------------------------------- the handle @@ -470,9 +623,10 @@ def __init__( """ iterable, specs = (args[0], args[1:]) if args and not isinstance(args[0], (str, dict)) else (None, args) self.smooth = smooth # default TWEMA weight for panels that don't set their own (wandb's smoothing slider) - self.panels = [self._with_defaults(_normalise_panel(p)) for p in specs] - self._explicit_layout = bool(self.panels) - self._placed = {n for p in self.panels for n in p["metrics"] + p["secondary"]} + self._panel_defaults: dict = {} # plot-level set_*() values, applied to panels created later + self._specs = [self._with_defaults(_normalise_panel(p)) for p in specs] + self._explicit_layout = bool(self._specs) + self._placed = {n for p in self._specs for n in p["metrics"] + p["secondary"]} self._layout = (max_cols, rows, cols, cell_size, dpi, unit) self._iterable, self._bar, self._desc, self._progress = iterable, None, desc, progress if total is None and iterable is not None: @@ -529,10 +683,72 @@ def _make_display_handle(): def _with_defaults(self, panel): if panel["smooth"] is None: panel["smooth"] = self.smooth + for key, value in self._panel_defaults.items(): # plot-level set_xlabel(...) etc. made before the panel existed + panel[key] = value return panel + # -- panels and axes, addressed like matplotlib ----------------------------------- + + @property + def panels(self) -> list: + """The panels of the grid, in order: `plot.panels[0].set_title("training")`.""" + return [Panel(self, i) for i in range(len(self._specs))] + + def __getitem__(self, metric: str) -> _Axis: + """ + The y-axis that holds `metric`: `plot["acc"].set_ylim(0, 1)`. Right-hand axes are found too, + which is what makes the left/right distinction disappear from the API. A metric that hasn't + been logged yet gets its panel created now. + """ + if metric not in self._placed: + self.data.setdefault(metric, ([], [])) + self._extend_layout([metric]) + for i, spec in enumerate(self._specs): + if metric in spec["metrics"]: + return _Axis(self, i, right=False) + if metric in spec["secondary"]: + return _Axis(self, i, right=True) + raise KeyError(metric) + + def _set_all(self, key, value): + """A plot-level setter: every existing panel, and every panel created later.""" + self._panel_defaults[key] = value + for spec in self._specs: + spec[key] = value + self._send_layout() + + def set_title(self, label): + self._set_all("title", str(label)) + + def set_xlabel(self, xlabel): + self._set_all("xlabel", str(xlabel)) + + def set_xlim(self, left=None, right=None): + self._set_all("xlim", _lim(left, right)) + + def set_ylabel(self, ylabel): + self._set_all("ylabel", str(ylabel)) + + def set_ylim(self, bottom=None, top=None): + self._set_all("ylim", _lim(bottom, top)) + + def set_yscale(self, value): + assert value in ("linear", "log"), f"yscale must be 'linear' or 'log', got {value!r}" + self._set_all("yscale", value) + + def set_smooth(self, weight): + """wandb-style smoothing weight in [0, 1) for every panel (and the default for later ones).""" + assert 0 <= weight < 1, f"smooth must be a weight in [0, 1), got {weight!r}" + self.smooth = weight + self._set_all("smooth", weight) + + def set(self, **kwargs): + """Like `Axes.set`, on every panel: `plot.set(xlabel="examples", yscale="log")`.""" + for key, value in kwargs.items(): + getattr(self, f"set_{key}")(value) + def _panels_or_placeholder(self): - return self.panels or [_normalise_panel({"title": "waiting for data…", "metrics": ["_"]})] + return self._specs or [_normalise_panel({"title": "waiting for data…", "metrics": ["_"]})] def _start_process(self): # "spawn", never "fork": the notebook process has usually initialised CUDA, @@ -680,12 +896,12 @@ def _extend_layout(self, new_metrics): return if self._explicit_layout: for m in unplaced: - self.panels.append(self._with_defaults(_normalise_panel({"metrics": [m]}))) - elif not self.panels: - self.panels.append(self._with_defaults(_normalise_panel({"metrics": unplaced}))) + self._specs.append(self._with_defaults(_normalise_panel({"metrics": [m]}))) + elif not self._specs: + self._specs.append(self._with_defaults(_normalise_panel({"metrics": unplaced}))) else: - self.panels[0]["metrics"].extend(unplaced) - self.panels[0]["title"] = " / ".join(self.panels[0]["metrics"]) + self._specs[0]["metrics"].extend(unplaced) + self._specs[0]["title"] = " / ".join(self._specs[0]["metrics"]) self._placed.update(unplaced) self._send_layout() @@ -696,15 +912,11 @@ def axhline(self, y, label=None, *, metric=None, **kwargs): artist. It goes on the axis of `metric`'s panel, or on the left axis of every panel if `metric` is None. E.g. the loss a model must beat: `plot.axhline(math.log(d_vocab), "uniform", metric="loss")`. """ + if metric is not None: + return self[metric].axhline(y, label, **kwargs) line = {**kwargs, "y": float(y), "label": str(label) if label is not None else f"{float(y):g}"} - if metric is not None and metric not in self._placed: - self.data.setdefault(metric, ([], [])) # create the metric's panel now so the line has somewhere to go - self._extend_layout([metric]) - for panel in self.panels: - if metric is None or metric in panel["metrics"]: - panel["axhlines"].append(line) - elif metric in panel["secondary"]: - panel["axhlines2"].append(line) + for spec in self._specs: + spec["axhlines"].append(line) self._send_layout() def axvline(self, x=None, label=None, **kwargs): @@ -722,9 +934,9 @@ def axvline(self, x=None, label=None, **kwargs): def _send_layout(self): if self.mode == "process": - self._inbox.put(("layout", self.panels)) + self._inbox.put(("layout", self._specs)) elif self.mode == "thread": - self._renderer.set_layout(self.panels) + self._renderer.set_layout(self._specs) def figure(self): """ diff --git a/mkdocs.yml b/mkdocs.yml index ff00d07..c0dc746 100644 --- a/mkdocs.yml +++ b/mkdocs.yml @@ -5,6 +5,7 @@ repo_url: https://github.com/ARENA-education/liveplot docs_dir: docs nav: - Home: index.md + - Guide: guide.md - API: api.md - Examples: examples.md theme: diff --git a/tests/test_axes_api.py b/tests/test_axes_api.py new file mode 100644 index 0000000..915d178 --- /dev/null +++ b/tests/test_axes_api.py @@ -0,0 +1,97 @@ +import pytest + +from liveplot import LivePlot +from liveplot.liveplot import Panel, _Axis + + +def test_addressing_by_metric_and_by_panel(): + p = LivePlot("loss | acc", "lr") + p.log(0, loss=1.0, acc=0.5, lr=1e-3) + assert [type(x) for x in p.panels] == [Panel, Panel] and p.panels[0].metrics == ["loss", "acc"] + left, right, lr = p["loss"], p["acc"], p["lr"] + assert isinstance(left, _Axis) and not left._right and right._right and left.panel.spec is p.panels[0].spec + assert right.metrics == ["acc"] and lr.panel.spec is p.panels[1].spec + assert repr(p.panels[0]) == "Panel(0: loss | acc)" + with pytest.raises(AssertionError): + p.panels[1].right # no metrics after '|' + + +def test_matplotlib_named_setters_go_to_the_right_axis(): + p = LivePlot("loss | acc") + p.log(0, loss=1.0, acc=0.5) + p["loss"].set_ylabel("cross-entropy") + p["acc"].set_ylabel("test accuracy") + p["acc"].set_ylim(0, 1) # matplotlib's two-argument form + p["loss"].set_ylim((0.0, 2.0)) # ... and the tuple form + p["acc"].set_yscale("log") + p.panels[0].set_title("training") + p.panels[0].set_xlabel("examples") + p.panels[0].set_xlim(0, 500) + spec = p.panels[0].spec + assert spec["ylabel"] == "cross-entropy" and spec["ylabel2"] == "test accuracy" + assert spec["ylim"] == (0.0, 2.0) and spec["ylim2"] == (0.0, 1.0) and spec["yscale2"] == "log" and spec["yscale"] == "linear" + assert spec["title"] == "training" and spec["xlabel"] == "examples" and spec["xlim"] == (0.0, 500.0) + p["acc"].set_title("via the axis, like a matplotlib Axes") # panel-level setter reachable from an axis + assert spec["title"] == "via the axis, like a matplotlib Axes" + p.panels[0].set_ylabel("left again") # y-setters on a panel act on its left axis + assert spec["ylabel"] == "left again" + with pytest.raises(AssertionError): + p["acc"].set_yscale("sqrt") + with pytest.raises(AssertionError): + p["acc"].set_ylim(0) + + +def test_set_kwargs_and_reference_lines(): + p = LivePlot("loss | acc") + p.log(0, loss=1.0, acc=0.5) + p["acc"].set(ylabel="acc", ylim=(0, 1), yscale="log") + p.panels[0].set(title="t", xlabel="x", smooth=0.5) + spec = p.panels[0].spec + assert (spec["ylabel2"], spec["ylim2"], spec["yscale2"], spec["title"], spec["xlabel"], spec["smooth"]) == ("acc", (0.0, 1.0), "log", "t", "x", 0.5) + p["loss"].axhline(0.1, "target", color="red") + p["acc"].axhline(0.9) + p.panels[0].axvline(3, "here", linewidth=2) + p.log(7, loss=0.5) + p.panels[0].axvline(label="now") # x defaults to the current step + assert spec["axhlines"] == [{"color": "red", "y": 0.1, "label": "target"}] and spec["axhlines2"] == [{"y": 0.9, "label": "0.9"}] + assert spec["axvlines"] == [{"linewidth": 2, "x": 3.0, "label": "here"}, {"x": 7.0, "label": "now"}] + assert p.axvlines == [], "panel-level axvline is not a plot-wide one" + p.axhline(0.0, "zero", metric="loss") # the plot-level form with metric= delegates to the axis + assert spec["axhlines"][-1] == {"y": 0.0, "label": "zero"} + + +def test_plot_level_setters_apply_to_every_panel_and_later_ones(): + p = LivePlot("loss", "acc") + p.set(xlabel="examples", yscale="log") + p.set_smooth(0.8) + p.log(0, loss=1.0, acc=0.5, lr=1e-3) # lr's panel is created after the setters were called + for panel in p.panels: + assert panel.spec["xlabel"] == "examples" and panel.spec["yscale"] == "log" and panel.spec["smooth"] == 0.8 + p.set_ylim(0, 10) + assert all(panel.spec["ylim"] == (0.0, 10.0) for panel in p.panels) + with pytest.raises(AssertionError): + p.set_smooth(1.5) + + +def test_zero_config_then_configure(): + p = LivePlot() + p.set_title("my run") # before any panel exists + p.log(0, loss=1.0) + assert p.panels[0].spec["title"] == "my run" + p["val_loss"].set_ylabel("held-out") # addressing a metric not logged yet creates its panel + assert p.panels[0].metrics == ["loss", "val_loss"], "no layout given: the new metric joins the shared panel" + assert p.panels[0].spec["ylabel"] == "held-out" + + +def test_figure_reflects_setters(): + p = LivePlot("loss | acc", progress=False) + for step in range(5): + p.log(step, loss=1.0 / (step + 1), acc=step / 5) + p["acc"].set_ylim(0, 1) + p["acc"].set_ylabel("accuracy") + p.panels[0].set_title("hello") + p.panels[0].axvline(2, "two") + fig = p.figure() + ax, ax2 = fig.axes[0], fig.axes[1] + assert ax.get_title() == "hello" and ax2.get_ylabel() == "accuracy" and ax2.get_ylim() == (0.0, 1.0) + assert any(t.get_text().strip() == "two" for t in ax.texts) diff --git a/tests/test_liveplot.py b/tests/test_liveplot.py index e38653c..5f7de7c 100644 --- a/tests/test_liveplot.py +++ b/tests/test_liveplot.py @@ -59,12 +59,12 @@ def test_off_mode_discovery_and_log_forms(): p.log(1, {"loss": 0.5}, acc=0.1) p.log(acc=0.2) # nothing wrapped: x stays at the last explicit step, 1 assert p.data == {"loss": ([0, 1], [1.0, 0.5]), "acc": ([1, 1], [0.1, 0.2])} - assert [pn["metrics"] for pn in p.panels] == [["loss", "acc"]], "no layout given: everything on one panel" + assert [pn["metrics"] for pn in p._specs] == [["loss", "acc"]], "no layout given: everything on one panel" p.finish(); p.finish() # idempotent q = LivePlot("loss", "a | b") q.log(0, loss=1, a=2, b=3, extra=4) - assert [(pn["metrics"], pn["secondary"]) for pn in q.panels] == [(["loss"], []), (["a"], ["b"]), (["extra"], [])] + assert [(pn["metrics"], pn["secondary"]) for pn in q._specs] == [(["loss"], []), (["a"], ["b"]), (["extra"], [])] def test_x_axis_follows_tqdm_counting(): @@ -165,7 +165,7 @@ def test_process_mode_end_to_end(fake_notebook): assert len(fake_notebook.frames) >= n_frames0 + 3 and all(f[:8] == PNG for f in fake_notebook.frames) assert not p._proc.is_alive() assert len(p.data["loss"][0]) == n_warmup + step and p.last_png == fake_notebook.frames[-1] - assert [pn["metrics"] for pn in p.panels] == [["loss"], ["acc"], ["lr"]] + assert [pn["metrics"] for pn in p._specs] == [["loss"], ["acc"], ["lr"]] def test_interrupt_inside_with_block_is_clean(fake_notebook): diff --git a/tests/test_reference_lines.py b/tests/test_reference_lines.py index 6e68772..ec46ca3 100644 --- a/tests/test_reference_lines.py +++ b/tests/test_reference_lines.py @@ -50,10 +50,10 @@ def test_axhline_and_axvline_off_mode(): p.axhline(0.2, "target", metric="loss", color="green") p.axhline(0.9, metric="acc") # label defaults to the value p.axhline(0.0) # no metric: left axis of every panel - assert p.panels[0]["axhlines"] == [{"color": "green", "y": 0.2, "label": "target"}, {"y": 0.0, "label": "0"}] - assert p.panels[0]["axhlines2"] == [{"y": 0.9, "label": "0.9"}] + assert p._specs[0]["axhlines"] == [{"color": "green", "y": 0.2, "label": "target"}, {"y": 0.0, "label": "0"}] + assert p._specs[0]["axhlines2"] == [{"y": 0.9, "label": "0.9"}] p.axhline(1e-3, "final lr", metric="lr") # metric not logged yet: creates its panel - assert p.panels[1]["metrics"] == ["lr"] and p.panels[1]["axhlines"] == [{"y": 0.001, "label": "final lr"}] + assert p._specs[1]["metrics"] == ["lr"] and p._specs[1]["axhlines"] == [{"y": 0.001, "label": "final lr"}] p.log(7, loss=0.5) p.axvline(label="lr drop") # x defaults to the current step p.axvline(9, "later", linewidth=2) diff --git a/tests/test_smoothing.py b/tests/test_smoothing.py index 7efa723..68879c8 100644 --- a/tests/test_smoothing.py +++ b/tests/test_smoothing.py @@ -77,7 +77,7 @@ def test_renderer_smooths_and_fades_raw(): def test_plot_wide_default_and_override(): p = LivePlot("loss", {"metrics": ["acc"], "smooth": 0}, smooth=0.9) p.log(0, loss=1.0, acc=0.5, lr=1e-3) # lr is discovered -> gets the default too - assert [pn["smooth"] for pn in p.panels] == [0.9, 0, 0.9] + assert [pn["smooth"] for pn in p._specs] == [0.9, 0, 0.9] q = LivePlot() q.log(0, loss=1.0) - assert q.panels[0]["smooth"] is None + assert q._specs[0]["smooth"] is None