diff --git a/README.md b/README.md index 06c54f5..10300fd 100644 --- a/README.md +++ b/README.md @@ -122,6 +122,96 @@ plot.axvline(label="lr drop") # every panel, at the plot.panels[1].axvline(2000, "checkpoint", linestyle="-") # one panel, at a given x ``` +## Images beside curves: `subplots` + +`LivePlot.subplots` mirrors `plt.subplots`, and a panel can hold a picture instead of lines. That +gives you the shape a GAN training loop wants: loss curves updating every step, generated samples +every so often, in one figure and one output cell. + +```python +plot, (ax_loss, ax_samples) = LivePlot.subplots(1, 2, total=epochs * len(loader), figsize=(11, 4)) +ax_loss.plot("lossD", "lossG") # which metrics live on this axis +ax_loss.twinx().plot("D(x)") # matplotlib's own spelling for a right-hand axis +ax_samples.set_title("generator samples") + +for epoch in range(epochs): + for imgs, _ in plot(loader, desc=f"epoch {epoch}"): + plot.log(lossD=..., lossG=...) # every step + if step % 250 == 0: + ax_samples.imshow(netG(fixed_noise), rows=2, vmin=-1, vmax=1) # replaces the last one +plot.finish() +``` + +`ax.plot("lossD", "lossG")` names metrics rather than passing data, which is matplotlib's own +`ax.plot("lossD", data=d)` form with the data source implicit -- the plot is the source, filled in +later by `log()`. `axes` follows matplotlib's squeeze rules: one panel for 1x1, a flat list for a +single row or column, a 2-d grid otherwise (`axes[0][1]` and `axes[0, 1]` both work, and +`axes.flat` walks it). `figsize` is the whole figure in inches, as matplotlib means it, and +`width_ratios` / `height_ratios` are matplotlib's too: `width_ratios=(1, 1.25)` gives a wide grid +of samples more room. `ax.legend(**kwargs)` goes to `Axes.legend`, so a crowded legend can move +off the curves: `ax.legend(loc="upper center", bbox_to_anchor=(0.5, -0.14), ncols=3)`. + +A training loop that already drives its own bar, one for the whole run rather than one per epoch, +keeps its shape: `plot.update()` and `plot.set_description(...)` are tqdm's manual-mode methods, +and the first `update()` opens the bar below the figure. + +```python +plot, (ax_loss, ax_samples) = LivePlot.subplots(1, 2, total=epochs * len(loader), width_ratios=(1, 1.25)) +for epoch in range(epochs): + plot.set_description(f"epoch {epoch}") + for imgs, _ in loader: + plot.log(lossD=..., lossG=...) + plot.update() +plot.finish() +``` + +[`examples/dcgan_synthetic.py`](examples/dcgan_synthetic.py) is the whole thing on a fake DCGAN run +laid out like ARENA's [0.5] trainer: smoothed losses with `ln 4` / `ln 2` reference lines and the +discriminator's outputs on a right-hand axis, beside ten CelebA faces coming out of the noise. + +For a single picture with no curves, `plot.imshow(x)` makes the panel on first use and every later +call replaces it: + +```python +plot = LivePlot() +for step in ...: + plot.imshow(model(holdout)) # overwrites, rather than stacking a new plot underneath +``` + +### What `imshow` accepts + +A batch is tiled into a grid for you: `rows=` or `cols=` alone infers the other, `griddim=(r, c)` +fixes both, padding with blanks or dropping the tail as needed. + +| input | read as | +|---|---| +| `(H, W)`, `(1, H, W)` | one grayscale image | +| `(H, W, 3)`, `(H, W, 4)` | one colour image, channels last | +| `(B, H, W)` | `B` grayscale images | +| `(B, 1\|3\|4, H, W)` | `B` images, channels first (torch's layout) | +| `(B, H, W, 1\|3\|4)` | `B` images, channels last | + +`(3, H, W)` and `(4, H, W)` are the ambiguous ones -- one colour image, or that many grayscale? -- +so they raise, naming the `channels=` to pass. matplotlib sidesteps this by refusing channels-first +outright; torch holds images that way, so the ambiguity is ours to resolve rather than ignore. + +Torch tensors go straight in: the conversion (`detach`, off the GPU, tile, scale to `uint8`) happens +on the calling thread, costs about 0.2 ms for ten 64x64 RGB samples, and means the render process +only ever unpickles a finished picture. Nothing in liveplot imports torch. + +### Scaling + +Values are scaled to the full range of the batch by default. `vmin` / `vmax` fix the range instead, +which is worth doing for a live view -- otherwise the black point moves every frame. `scale_each=True` +scales each image on its own, as `torchvision.utils.make_grid` does; it is off by default because it +hides exactly what you watch samples for, a washed-out or collapsed one stops looking anomalous. +`uint8` input passes through untouched. + +Note that liveplot does this scaling itself rather than leaving it to matplotlib, and so `vmin` / +`vmax` work for colour images too. matplotlib's `imshow` **ignores** them for RGB(A) data and merely +clips floats to `[0, 1]`: a generator ending in `tanh` outputs `[-1, 1]`, and roughly 45% of it would +go to black with no way to say otherwise. + ## 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. diff --git a/docs/imshow.gif b/docs/imshow.gif new file mode 100644 index 0000000..1608b36 Binary files /dev/null and b/docs/imshow.gif differ diff --git a/examples/dcgan_synthetic.py b/examples/dcgan_synthetic.py new file mode 100644 index 0000000..2afd3f0 --- /dev/null +++ b/examples/dcgan_synthetic.py @@ -0,0 +1,164 @@ +# %% +""" +A synthetic DCGAN run, laid out like `DCGANTrainer` in ARENA's [0.5] VAEs & GANs day, to show what +liveplot looks like on that training loop without a GPU. Nothing is trained: the losses are made up +to behave the way the course describes (lossD starting at ln 4 and falling, lossG starting at ln 2 +and rising, D(x) and D(G(z)) pulling apart from 1/2), and the "generator" returns nine real CelebA +faces behind noise that fades as training goes on. + +Run it cell by cell in the VS Code / Cursor interactive window or Jupyter to see it live: the loss +curves and the samples side by side in one output, with a tqdm bar underneath. Run it as a script +(`python examples/dcgan_synthetic.py`) and nothing is drawn on screen, but it records the run to +`dcgan_synthetic.gif`. + +Needs torch and Pillow, and fetches nine CelebA faces (about 60 kB, cached in ~/.cache/liveplot) +through Hugging Face's dataset viewer, not the 1.4 GB dataset. + +The parts to copy into the real trainer are marked `# liveplot:`. +""" + +import io +import json +import math +import sys +import time +import urllib.request +from dataclasses import dataclass +from pathlib import Path + +import numpy as np +import torch as t +from PIL import Image + +from liveplot import LivePlot + +MAIN = __name__ == "__main__" +CELEB_IMAGE_SIZE = 64 + + +# %% +# The data: a few CelebA faces, prepared as the course's `get_dataset("CELEB")` does (resize the short +# side to 64, centre-crop to 64x64, uint8), then mapped to [-1, 1] like TANH_RANGE_TRANSFORM. + + +def celeba_faces(n: int = 10) -> t.Tensor: + """(n, 3, 64, 64) uint8: the first n images of nielsr/CelebA-faces, the dataset the course uses.""" + cache = Path.home() / ".cache" / "liveplot" / f"celeba_faces_{n}.npy" + if not cache.exists(): + url = f"https://datasets-server.huggingface.co/rows?dataset=nielsr/CelebA-faces&config=default&split=train&offset=0&length={n}" + rows = json.load(urllib.request.urlopen(url, timeout=60))["rows"] + faces = [] + for row in rows: + img = Image.open(io.BytesIO(urllib.request.urlopen(row["row"]["image"]["src"], timeout=60).read())).convert("RGB") + scale = CELEB_IMAGE_SIZE / min(img.size) # transforms.Resize(64): the short side to 64 + img = img.resize((round(img.width * scale), round(img.height * scale)), Image.BILINEAR) + left, top = (img.width - CELEB_IMAGE_SIZE) // 2, (img.height - CELEB_IMAGE_SIZE) // 2 + img = img.crop((left, top, left + CELEB_IMAGE_SIZE, top + CELEB_IMAGE_SIZE)) # CenterCrop(64) + faces.append(np.asarray(img).transpose(2, 0, 1)) + cache.parent.mkdir(parents=True, exist_ok=True) + np.save(cache, np.stack(faces)) + return t.from_numpy(np.load(cache)) + + +def tanh_range(x: t.Tensor) -> t.Tensor: + """TANH_RANGE_TRANSFORM on a uint8 batch: [0, 255] -> [-1, 1].""" + return x.float() / 127.5 - 1 + + +# %% +# The trainer. Same shape as the course's DCGANTrainer; the two training steps and the generator are +# fakes, and everything liveplot adds is marked. + + +@dataclass +class DCGANArgs: + batch_size: int = 64 + epochs: int = 3 + batches_per_epoch: int = 150 # stands in for len(self.trainloader) + log_every_n_steps: int = 15 # 250 in the course, with ~3000 batches per epoch + seconds_per_step: float = 0.03 # stands in for the GPU's time per batch + seed: int = 0 + + +class SyntheticDCGANTrainer: + def __init__(self, args: DCGANArgs): + self.args = args + self.gen = t.Generator().manual_seed(args.seed) + self.faces = tanh_range(celeba_faces(9)) # what a well-trained netG(self.fixed_noise) would give + self.fixed_noise = t.randn(self.faces.shape, generator=self.gen) # the fake generator's own noise + self.total_steps = args.epochs * args.batches_per_epoch + self.d_state = t.zeros(2) # slowly wandering offsets that make D(x), D(G(z)) look like a real run + + def _progress(self) -> float: + return self.step / self.total_steps + + def training_step_discriminator(self) -> tuple[float, float, float]: + """Returns lossD and the discriminator's mean outputs D(x) (real) and D(G(z)) (fake).""" + gap = 0.26 * (1 - math.exp(-6 * self._progress())) # D learns to separate real from fake, fast then slow + self.d_state = 0.9 * self.d_state + 0.04 * t.randn(2, generator=self.gen) + spike = 0.12 * float(t.rand(1, generator=self.gen) < 0.01) # the odd batch D gets badly wrong + d_real = min(max(0.5 + gap + float(self.d_state[0]) - spike, 0.02), 0.98) + d_fake = min(max(0.5 - gap + float(self.d_state[1]) + spike, 0.02), 0.98) + lossD = -(math.log(d_real) + math.log(1 - d_fake)) # the course's lossD, at the batch-mean outputs + return lossD, d_real, d_fake + + def training_step_generator(self, d_fake: float) -> float: + """lossG = -log D(G(z)), after the discriminator's step has moved D(G(z)) a little.""" + d_fake = min(max(d_fake + 0.03 * float(t.randn(1, generator=self.gen)), 0.02), 0.98) + return -math.log(d_fake) + + @t.inference_mode() + def log_samples(self) -> None: + """netG(self.fixed_noise), faked: the faces, behind noise that fades as training goes on.""" + quality = 1 - math.exp(-4 * self._progress()) + output = quality * self.faces + (1 - quality) * self.fixed_noise + # Clip values to make the visualization clearer (as the course does) + output = output.clamp(output.quantile(0.01), output.quantile(0.99)) + self.ax_samples.imshow(output) # liveplot: replaces LiveImage.update / wandb.Image (9 -> a 3x3 grid) + + def train(self, **plot_kwargs) -> LivePlot: + self.step = 0 + # liveplot: one figure for the curves and the samples, and the tqdm bar underneath it (this + # replaces `self.live_image = LiveImage()` and `progress_bar = tqdm(total=...)`) + self.plot, (ax_loss, self.ax_samples) = LivePlot.subplots( + 1, 2, total=self.total_steps, figsize=(12, 4.5), **plot_kwargs + ) + ax_loss.plot("lossD", "lossG") + ax_loss.set_smooth(0.6) # per-batch GAN losses are noisy: wandb's smoothing, raw values faded behind + ax_loss.axhline(math.log(4), "ln 4: lossD, D at chance") # where lossD starts, and where a perfect G ends + ax_loss.axhline(math.log(2), "ln 2: lossG, D at chance", linestyle=":") # ... and the same for lossG + ax_loss.twinx().plot("D(x)", "D(G(z))").set_ylim(0, 1) + ax_loss.set_title("losses (left), discriminator outputs (right)") + ax_loss.legend(loc="upper center", bbox_to_anchor=(0.5, -0.14), ncols=3) # under the panel, off the curves + self.ax_samples.set_title("netG(fixed_noise)") + + for epoch in range(self.args.epochs): + self.plot.set_description(f"epoch {epoch}") # liveplot: was progress_bar.set_description + for _ in range(self.args.batches_per_epoch): # for img_real, label in self.trainloader: + lossD, d_real, d_fake = self.training_step_discriminator() + lossG = self.training_step_generator(d_fake) + + # liveplot: log the step's numbers and advance the bar (was progress_bar.update()) + self.plot.log({"lossD": lossD, "lossG": lossG, "D(x)": d_real, "D(G(z))": d_fake}) + self.plot.update() + self.step += 1 + + if self.step % self.args.log_every_n_steps == 0: + self.log_samples() + time.sleep(self.args.seconds_per_step) + + self.log_samples() # the final samples, whatever the step count + self.plot.finish() + return self.plot + + +# %% +# Run it. In a notebook or interactive window you get the live figure with a bar underneath; as a +# script, the run is recorded to a GIF instead (pass a path to choose where). + +if MAIN: + in_notebook = "ipykernel" in sys.modules + gif = None if in_notebook else (sys.argv[1] if len(sys.argv) > 1 else "dcgan_synthetic.gif") + plot = SyntheticDCGANTrainer(DCGANArgs()).train(**({"record": gif, "refresh_seconds": 0.25} if gif else {})) + if gif: + print(f"{len(plot.frames)} frames -> {gif}") diff --git a/examples/demo.py b/examples/demo.py index 05f7a17..871f222 100644 --- a/examples/demo.py +++ b/examples/demo.py @@ -179,3 +179,59 @@ def train_classifier(steps=600, batch_size=64, lr=0.5, eval_every=50, seed=0, pa if MAIN: xs, losses = plot.data["loss"] print(f"{len(xs)} loss points, x from {xs[0]} to {xs[-1]} examples; final eval acc {plot.latest['eval_acc']:.3f}") + +# %% +# 9. Images beside curves: LivePlot.subplots mirrors plt.subplots, and a panel can hold a picture +# instead of lines. `ax.plot("name", ...)` says which metrics live on an axis (matplotlib spells +# the same idea `ax.plot("name", data=d)`); `ax.imshow(tensor)` draws a batch as a grid and +# replaces it on every call. This is the shape a GAN training loop wants: loss curves updating +# every step, generated samples every so often, in one figure. +# +# A batch is tiled for you: `rows=` or `cols=` infers the other, and values are scaled to the +# full range unless you fix it with vmin/vmax -- worth doing here, since a generator that ends +# in tanh always produces [-1, 1] and a fixed range keeps the black point still between frames. + + +def fake_generator(step, rng, n=8, size=24): + """Stand-in for a generator: blobs that sharpen as `step` grows, in tanh range like a real one.""" + y, x = np.mgrid[0:size, 0:size] / size + out = np.empty((n, 3, size, size)) + for i in range(n): + cx, cy, sharp = rng.random(), rng.random(), 0.02 + 0.5 / (1 + step / 40) + blob = np.exp(-((x - cx) ** 2 + (y - cy) ** 2) / sharp) + out[i] = np.stack([blob, blob * (0.3 + 0.7 * cx), blob * (0.3 + 0.7 * cy)]) + return out * 2 - 1 # [-1, 1], as a tanh output would be + + +if MAIN: + rng = np.random.default_rng(0) + plot, (ax_loss, ax_samples) = LivePlot.subplots(1, 2, total=400, figsize=(11, 4)) + ax_loss.plot("lossD", "lossG") + ax_loss.twinx().plot("D(x)") # matplotlib's own spelling for a right-hand axis + ax_samples.set_title("generator samples") + + for step in plot(range(400)): + plot.log(lossD=1 / (1 + step / 50) + 0.05 * rng.random(), + lossG=0.4 + 0.3 * rng.random(), + **{"D(x)": 0.5 + 0.3 / (1 + step / 80)}) + if step % 25 == 0: + ax_samples.imshow(fake_generator(step, rng), rows=2, vmin=-1, vmax=1) + slow() + plot.finish() + +# %% +# 10. One image on its own, overwritten in place: plot.imshow() makes the panel on first use. The +# accepted layouts are (H, W), (C, H, W), (H, W, C), (B, H, W), (B, C, H, W) and (B, H, W, C). +# (3, H, W) and (4, H, W) are the ambiguous ones -- one colour image, or that many grayscale? -- +# and they raise, telling you which `channels=` to pass. + +if MAIN: + plot = LivePlot() + for step in range(40): + plot.imshow(np.random.default_rng(step).random((6, 1, 28, 28)), rows=2) # replaces the last + slow(0.05) + plot.finish() + try: + plot.imshow(np.zeros((3, 8, 8))) + except ValueError as e: + print(e) diff --git a/liveplot/_images.py b/liveplot/_images.py new file mode 100644 index 0000000..94a48b2 --- /dev/null +++ b/liveplot/_images.py @@ -0,0 +1,189 @@ +""" +Turning a tensor into something matplotlib will draw: reading its layout, tiling a batch into a +grid, and scaling the values to uint8. + +Nothing here imports torch, and numpy only inside the functions that need it, so `import liveplot` +stays free (see the note at the top of liveplot.py). A caller may hand us a torch tensor, a numpy +array, or anything else with `__array__`. + +Why we scale the values ourselves instead of letting `imshow` do it: matplotlib normalises *scalar* +data with vmin/vmax, but for RGB(A) data it ignores vmin/vmax completely and simply clips floats to +[0, 1] -- checked on matplotlib 3.9 and 3.11, where `imshow(rgb, vmin=-1, vmax=1)` renders pixel for +pixel the same as `imshow(clip(rgb, 0, 1))`. A DCGAN generator ends in tanh, so its output lives in +[-1, 1] and matplotlib would clip roughly 45% of it to black with no way to say otherwise. So we do +the scaling and hand `imshow` uint8, which it passes through untouched. +""" + +from __future__ import annotations + +import math + +_CHANNELS = (1, 3, 4) # channel counts we recognise: grayscale, RGB, RGBA + + +def _as_array(x): + """Anything array-like -> a numpy array, without naming torch.""" + import numpy as np + + if hasattr(x, "detach"): + x = x.detach() # a tensor that wants grad + if hasattr(x, "cpu"): + x = x.cpu() # ... and may be on a GPU, where numpy can't see it + if getattr(x, "is_floating_point", lambda: False)(): + x = x.float() # bfloat16 (autocast output) has no numpy dtype; float16 is widened too, harmlessly + return np.asarray(x) + + +def _ambiguous(shape, n): + kind = "RGBA" if n == 4 else "RGB" + return ValueError( + f"ambiguous shape {shape}: one {kind} image, or {n} grayscale images?\n" + f' channels="first" -> one {kind} image (or pass shape {(1,) + shape})\n' + f' channels="none" -> {n} grayscale images (or pass shape {(n, 1) + shape[1:]})' + ) + + +def _as_bhwc(a, channels: str | None): + """ + Any accepted layout -> (batch, height, width, channels). + + `channels` says how to read a 3-d input: "first" for (C, H, W), "last" for (H, W, C), "none" + for a batch of grayscale (B, H, W). None works it out: + + (H, W) one grayscale image + (1, H, W) one grayscale image -- the two readings draw the same pixels + (3, H, W) / (4, H, W) ambiguous, raises (see _ambiguous) + (H, W, 3) / (H, W, 4) one image, channels last + (B, H, W) B grayscale images + (B, 1|3|4, H, W) B images, channels first + (B, H, W, 1|3|4) B images, channels last + + matplotlib itself sidesteps all of this by refusing channels-first outright ("Invalid shape + (3, H, W) for image data"); it only ever takes (H, W), (H, W, 1), (H, W, 3) or (H, W, 4). We + have to accept channels-first because that is how torch holds images, so the ambiguity it + designed away is ours to resolve. + """ + import numpy as np + + assert channels in (None, "first", "last", "none"), \ + f'channels must be "first", "last", "none" or None, got {channels!r}' + shape = tuple(a.shape) + + if a.ndim == 2: + return a[None, :, :, None] + + if a.ndim == 3: + if channels == "first": + assert shape[0] in _CHANNELS, f'channels="first" needs {_CHANNELS} channels, got shape {shape}' + return np.moveaxis(a, 0, -1)[None] + if channels == "last": + assert shape[-1] in _CHANNELS, f'channels="last" needs {_CHANNELS} channels, got shape {shape}' + return a[None] + if channels == "none": + return a[..., None] + if shape[0] == 1: # (1, H, W): one image either way + return np.moveaxis(a, 0, -1)[None] + if shape[0] in (3, 4): + raise _ambiguous(shape, shape[0]) + if shape[-1] in (3, 4): # a batch of grayscale 3 or 4 pixels wide is not a thing + return a[None] + return a[..., None] + + if a.ndim == 4: + assert channels != "none", 'channels="none" describes a 3-d (B, H, W) input, not a 4-d one' + if channels == "last" or (channels is None and shape[1] not in _CHANNELS): + assert shape[-1] in _CHANNELS, ( + f"shape {shape} is not a batch of images: expected {_CHANNELS} channels at dim 1 " + f"(B, C, H, W) or dim 3 (B, H, W, C)" + ) + return a + return np.moveaxis(a, 1, -1) # (B, C, H, W), torch's layout, preferred when both could fit + + raise ValueError(f"imshow needs a 2-, 3- or 4-d array, got shape {shape}") + + +def _to_uint8(a, vmin, vmax, scale_each: bool): + """Scale to 0..255. uint8 passes straight through unless vmin/vmax ask for a rescale.""" + import numpy as np + + if a.dtype == np.uint8 and vmin is None and vmax is None: + return a + a = a.astype(np.float32, copy=False) + if scale_each: # per image, like torchvision's make_grid(scale_each=True) + lo = a.min(axis=(1, 2, 3), keepdims=True) if vmin is None else np.float32(vmin) + hi = a.max(axis=(1, 2, 3), keepdims=True) if vmax is None else np.float32(vmax) + else: + lo = np.float32(a.min() if vmin is None else vmin) + hi = np.float32(a.max() if vmax is None else vmax) + a = (a - lo) / np.maximum(hi - lo, 1e-12) # a constant image comes out at 0, as in make_grid + return (np.clip(a, 0.0, 1.0) * 255).round().astype(np.uint8) + + +def grid_shape(n: int, rows: int | None, cols: int | None, griddim=None, max_cols: int | None = 8) -> tuple[int, int]: + """ + Rows x cols for `n` images. `griddim=(rows, cols)` fixes both; one of `rows` / `cols` infers + the other; with neither, the near-square rule liveplot already uses for its panel grid (3 -> 2x2, + 5 -> 3x2, 8 -> 3x3, 64 -> 8x8), at most `max_cols` wide. Unlike the panel grid, the result need + not fit: the caller pads short and drops long. + """ + if griddim is not None: + assert rows is None and cols is None, "give griddim, or rows/cols, not both" + rows, cols = griddim + if rows is None and cols is None: + rows = math.ceil(math.sqrt(n)) + cols = math.ceil(n / rows) + if max_cols is not None and cols > max_cols: + cols = max_cols + rows = math.ceil(n / cols) + elif rows is None: + rows = math.ceil(n / cols) + elif cols is None: + cols = math.ceil(n / rows) + rows, cols = int(rows), int(cols) + assert rows >= 1 and cols >= 1, f"grid must be at least 1x1, got {rows}x{cols}" + return rows, cols + + +def _tile(a, rows: int, cols: int, pad_value: int, padding: int): + """ + (B, H, W, C) -> one (H', W', C) grid with `padding` pixels of `pad_value` around every image, as + torchvision's make_grid lays it out; missing images are blank cells, extra ones are dropped. A + single image gets no border. + """ + import numpy as np + + _, h, w, c = a.shape + p = padding if rows * cols > 1 else 0 + out = np.full((rows * (h + p) + p, cols * (w + p) + p, c), pad_value, dtype=a.dtype) + for k, img in enumerate(a[: rows * cols]): # more images than cells: the rest fall off the end + r, col = divmod(k, cols) + out[p + r * (h + p): p + r * (h + p) + h, p + col * (w + p): p + col * (w + p) + w] = img + return out + + +def to_grid(x, *, rows=None, cols=None, griddim=None, vmin=None, vmax=None, scale_each=False, + channels=None, padding=0, pad_value=0, max_images: int | None = 64, max_cols: int | None = 8): + """ + A tensor -> one uint8 image, (H, W) for grayscale or (H, W, 3|4), ready for `Axes.imshow`. + Runs on the calling thread: tiling 10 64x64 RGB samples costs ~0.2 ms and quarters what the + render process has to unpickle, and a CUDA tensor has to come back to the host here anyway. + + Only the first `max_images` of a batch are shown (with a warning), since a grid of 1000 + thumbnails is rarely what anyone meant; `max_images=None` shows them all. Images sit edge to + edge; `padding` pixels of `pad_value` between them if you want gaps, as make_grid draws them. + """ + import warnings + + a = _as_bhwc(_as_array(x), channels) + if max_images is not None and a.shape[0] > max_images: + warnings.warn( + f"imshow: showing the first {max_images} of {a.shape[0]} images " + f"(slice the batch yourself, or pass max_images=None to show them all)", + stacklevel=3, + ) + a = a[:max_images] + rows, cols = grid_shape(a.shape[0], rows, cols, griddim, max_cols) + a = a[: rows * cols] # drop before scaling: an image nobody can see must not set the range + a = _to_uint8(a, vmin, vmax, scale_each) + grid = _tile(a, rows, cols, pad_value, padding) + return grid[..., 0] if grid.shape[-1] == 1 else grid diff --git a/liveplot/liveplot.py b/liveplot/liveplot.py index 729bff7..8da4554 100644 --- a/liveplot/liveplot.py +++ b/liveplot/liveplot.py @@ -64,6 +64,26 @@ time-weighted EMA with the same 0 to 1 weight (`smooth=0.9`), raw values faded behind. +Images beside curves: `LivePlot.subplots` mirrors `plt.subplots`, and a panel can +hold a picture instead of lines -- loss curves updating every step next to samples +from a generator, in one figure: + + plot, (ax_loss, ax_samples) = LivePlot.subplots(1, 2, total=n_steps, figsize=(11, 4)) + ax_loss.plot("lossD", "lossG") # which metrics live on this axis + ax_loss.twinx().plot("D(x)") # matplotlib's spelling for a right-hand axis + ... + ax_samples.imshow(netG(noise), rows=2, vmin=-1, vmax=1) # replaces the last one + +`plot.imshow(x)` does the same on a plot of its own, making the panel on first use. +A batch is tiled for you (`rows` / `cols` / `griddim=(r, c)`, padding or dropping +to fit, with make_grid's 2-pixel gaps); (H, W), (C, H, W), (H, W, C), (B, H, W), (B, C, H, W) and (B, H, W, C) are +all understood, and the two genuinely ambiguous shapes, (3, H, W) and (4, H, W), +raise and name the `channels=` to pass. Values are scaled to the batch's full range +unless `vmin` / `vmax` fix it (worth doing live: otherwise the black point moves +every frame), or `scale_each=True` scales each image alone, as make_grid does. +Unlike matplotlib, which ignores vmin/vmax for colour data and clips it to [0, 1], +these apply to colour images too. + 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 `refresh_seconds` (default 1.0; points arriving in between are batched into the @@ -112,7 +132,7 @@ import weakref _PANEL_KEYS = {"title", "metrics", "secondary", "xlabel", "ylabel", "ylabel2", "xlim", "ylim", "ylim2", "axhlines", "axhlines2", - "axvlines", "smooth", "yscale", "yscale2"} + "axvlines", "smooth", "yscale", "yscale2", "kind", "cmap", "legend"} _REF_LINE_STYLE = {"linestyle": "--", "linewidth": 1, "color": "0.45"} # defaults for axhline / axvline artists @@ -153,13 +173,16 @@ def _check_smooth(weight): return weight -def _normalise_panel(spec) -> dict: +def _normalise_panel(spec, allow_empty: bool = False) -> dict: + """`allow_empty` is for the panels `LivePlot.subplots` hands out before anything is drawn on them.""" if isinstance(spec, str): spec = _parse_panel_string(spec) unknown = set(spec) - _PANEL_KEYS assert not unknown, f"unknown panel keys {sorted(unknown)}; allowed: {sorted(_PANEL_KEYS)}" + kind = spec.get("kind", "curve") + assert kind in ("curve", "image"), f'kind must be "curve" or "image", got {kind!r}' metrics, secondary = list(spec.get("metrics", [])), list(spec.get("secondary", [])) - assert metrics or secondary, f"a panel needs at least one metric: {spec!r}" + assert metrics or secondary or kind == "image" or allow_empty, f"a panel needs at least one metric: {spec!r}" for lim in ("xlim", "ylim", "ylim2"): if spec.get(lim) is not None: assert len(spec[lim]) == 2, f"{lim} must be a (low, high) pair" @@ -179,6 +202,9 @@ def _normalise_panel(spec) -> dict: "smooth": _check_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 + "kind": kind, # "curve" (metric lines) or "image" (one picture, overwritten in place) + "cmap": spec.get("cmap", "gray"), # image panels only; matplotlib ignores it for RGB data + "legend": dict(spec.get("legend") or {}), # Axes.legend kwargs (loc, ncols, bbox_to_anchor, ...) } @@ -239,9 +265,10 @@ class _FigureRenderer: rebuilds the figure for a new panel list, keeping the history. """ - def __init__(self, panels, xlim, layout, hist=None): + def __init__(self, panels, xlim, layout, hist=None, images=None): self.xlim, self.layout = xlim, layout # xlim = default x range or None; layout = (max_cols, rows, cols, cell_size, dpi, xlabel) self.hist = hist if hist is not None else {} + self.images = dict(images) if images else {} # panel index -> (the uint8 array it is showing, its x) self.axvlines: list[dict] = [] # vertical reference lines (matplotlib axvline kwargs), drawn on every panel self.set_layout(panels) @@ -250,19 +277,38 @@ def set_layout(self, panels): from matplotlib.backends.backend_agg import FigureCanvasAgg from matplotlib.figure import Figure - max_cols, rows_opt, cols_opt, cell_size, dpi, xlabel = self.layout + max_cols, rows_opt, cols_opt, cell_size, dpi, xlabel, *rest = self.layout + gridspec_kw = rest[0] if rest else {} # width_ratios / height_ratios, from subplots() n = len(panels) rows, cols = _grid_shape(n, max_cols, rows_opt, cols_opt) - self.fig = Figure(figsize=(cell_size[0] * cols, cell_size[1] * rows)) + self.panels = panels + fig_width = cell_size[0] * cols + # One row with pictures in it, and no width_ratios chosen: size each image panel to its picture + # at the row's height, and give the curves what is left, so neither is framed in white space. + self.fitted_aspects = None + if rows == 1 and "width_ratios" not in gridspec_kw and any(p["kind"] == "image" for p in panels): + widths, self.fitted_aspects = self._fitted_widths(panels, cell_size, fig_width) + fig_width = sum(widths) + gridspec_kw = {**gridspec_kw, "width_ratios": widths + [widths[-1]] * (cols - n)} + self.fig = Figure(figsize=(fig_width, cell_size[1] * rows)) FigureCanvasAgg(self.fig) self.dpi = dpi - axes = self.fig.subplots(rows, cols, squeeze=False).flatten() + axes = self.fig.subplots(rows, cols, squeeze=False, gridspec_kw=gridspec_kw).flatten() palette = matplotlib.rcParams["axes.prop_cycle"].by_key()["color"] self.lines = {} self.raw_lines = {} # metric -> faded line of the unsmoothed values, on smoothed panels self.smooth = {} # metric -> TWEMA weight in (0, 1) or None self.fixed = {} # axes -> (x fixed?, y fixed?): fixed axes are never autoscaled - for ax, panel in zip(axes, panels): + self.panel_cmaps = {i: p["cmap"] for i, p in enumerate(panels)} + self.panel_titles = {i: p["title"] for i, p in enumerate(panels)} + self.image_axes = {} # panel index -> its Axes, for the image panels + self.image_artists = {} # ... and the AxesImage drawn on it, so a redraw can set_data in place + for i, (ax, panel) in enumerate(zip(axes, panels)): + if panel["kind"] == "image": + ax.set_title(panel["title"]) + ax.set_axis_off() # pixel indices along the edge of a tiled grid are just noise + self.image_axes[i] = ax + continue names = panel["metrics"] + panel["secondary"] ax2 = ax.twinx() if panel["secondary"] else None ax.set_yscale(panel["yscale"]) @@ -293,8 +339,9 @@ def set_layout(self, panels): if panel["ylim"]: ax.set_ylim(*panel["ylim"]) self.fixed[ax] = (xlim is not None, panel["ylim"] is not None) - if names != ["_"]: # (the "waiting for data" placeholder has no legend) - ax.legend(handles, names, loc="best", fontsize=8) # every panel gets a legend (reference lines included) + if names and names != ["_"]: # (nothing logged yet, or the "waiting for data" placeholder) + # every panel gets a legend (reference lines included); Panel.legend(**kwargs) moves or styles it + ax.legend(handles, names, **{"loc": "best", "fontsize": 8, **panel["legend"]}) if ax2 is not None: ax.set_ylabel(panel["ylabel"] or ", ".join(panel["metrics"])) ax2.set_ylabel(panel["ylabel2"] or ", ".join(panel["secondary"])) @@ -311,8 +358,32 @@ def set_layout(self, panels): self._draw_axvline(kw, [ax]) for kw in self.axvlines: self._draw_axvline(kw) - self.fig.tight_layout() + for index in self.images: # a re-layout rebuilds the figure; put the pictures back + if index in self.image_axes: + self._draw_image(index) + self.fig.tight_layout(pad=0.6) + def _fitted_widths(self, panels, cell_size, fig_width): + """ + Panel widths in inches for a one-row figure: an image panel is as wide as its picture is at the + row's height (less room for the title), the curve panels share the rest of the figure's width, + or keep their own if the pictures leave too little. Also returns the aspect ratio each image + panel was sized for, so a picture of another shape can trigger a re-layout. + """ + cell_w, cell_h = cell_size + aspects = {} + for i, panel in enumerate(panels): + if panel["kind"] == "image": + arr = self.images.get(i, (None,))[0] + aspects[i] = None if arr is None else arr.shape[1] / arr.shape[0] + image_h = cell_h - 0.55 # the title above, and tight_layout's margins + widths = [0.2 + image_h * aspects[i] if aspects.get(i) else cell_w for i in range(len(panels))] + curves = [i for i, panel in enumerate(panels) if panel["kind"] != "image"] + if curves: + spare = (fig_width - sum(widths[i] for i in aspects)) / len(curves) + for i in curves: + widths[i] = max(spare, 0.6 * cell_w) + return widths, aspects def add_axvline(self, kw: dict): self.axvlines.append(kw) self._draw_axvline(kw) @@ -320,11 +391,34 @@ def add_axvline(self, kw: dict): 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 if axes is None else axes: + picture = set(self.image_axes.values()) # a vertical line across a picture means nothing + for ax in ([a for a in self.axes if a not in picture] 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) + def add_image(self, index: int, arr, x=None): + self.images[index] = (arr, x) + if index in self.image_axes: + self._draw_image(index) + + def _draw_image(self, index: int): + arr, x = self.images[index] + if self.fitted_aspects is not None and self.fitted_aspects.get(index) != arr.shape[1] / arr.shape[0]: + return self.set_layout(self.panels) # a picture of a new shape: re-fit the panel widths (draws it too) + ax = self.image_axes[index] + if x is not None: # say how old the picture is: it is only redrawn when imshow is called + title, when = self.panel_titles[index], f"{self.layout[5]} {x:g}" + ax.set_title(f"{title} ({when})" if title else when) + artist = self.image_artists.get(index) + if artist is not None and artist.get_array().shape == arr.shape: + artist.set_data(arr) # same size as last time: no new artist, no rescale + return + if artist is not None: + artist.remove() + cmap = self.panel_cmaps.get(index, "gray") + self.image_artists[index] = ax.imshow(arr, cmap=cmap, interpolation="nearest") + def add(self, step, metrics: dict): for name, value in metrics.items(): xs, ys = self.hist.setdefault(name, ([], [])) @@ -356,7 +450,7 @@ def _render_worker(panels, inbox, outbox, refresh_seconds, layout, xlim, parent_ """ Loop: collect messages from `inbox`, redraw at most once per `refresh_seconds` while there is new data, put PNG bytes on `outbox`. Messages: ("data", step, - metrics); ("layout", panels) to rebuild the figure; None to finish (draw one + metrics); ("layout", panels, layout) to rebuild the figure; None to finish (draw one last frame, put None, exit). Exits on its own if the parent process is gone. """ signal.signal(signal.SIGINT, signal.SIG_IGN) # Jupyter interrupts the whole process group; not our business @@ -388,11 +482,15 @@ def _render_worker(panels, inbox, outbox, refresh_seconds, layout, xlim, parent_ if item is None: running = False elif item[0] == "layout": + renderer.layout = item[2] # grid options can change after start (subplots' width_ratios) renderer.set_layout(item[1]) dirty = True elif item[0] == "axvline": renderer.add_axvline(item[1]) dirty = True + elif item[0] == "image": + renderer.add_image(*item[1:]) + dirty = True else: renderer.add(item[1], item[2]) dirty = True @@ -483,6 +581,26 @@ def _set(self, key, value): self._plot._specs[self._index][key + ("2" if self._right else "")] = value self._plot._send_layout() + def plot(self, *metrics): + """ + Put these metrics on this axis: `ax.plot("lossD", "lossG")`. matplotlib spells the same + thing `ax.plot("lossD", data=d)`, naming series in a data source; here the source is the + plot itself, filled in later by `log()`. Returns the axis, so calls can be chained. + """ + assert metrics, "plot() needs at least one metric name" + assert all(isinstance(m, str) for m in metrics), \ + f"plot() takes metric names, not data: {[m for m in metrics if not isinstance(m, str)]!r}" + panel = self._plot._specs[self._index] + assert panel["kind"] == "curve", "this panel is showing an image; use a different panel for curves" + key = "secondary" if self._right else "metrics" + auto = " / ".join(panel["metrics"] + panel["secondary"]) + was_auto = panel["title"] in (auto, "") # unnamed, or named after the metrics it had + panel[key].extend(m for m in metrics if m not in panel[key]) + if was_auto: # a title set with set_title() survives a later plot() call + panel["title"] = " / ".join(panel["metrics"] + panel["secondary"]) + self._plot._place(metrics) + return self + def set_ylabel(self, ylabel): self._set("ylabel", str(ylabel)) @@ -508,6 +626,9 @@ def set(self, **kwargs): def set_title(self, label): self.panel.set_title(label) + def legend(self, **kwargs): + self.panel.legend(**kwargs) + def set_xlabel(self, xlabel): self.panel.set_xlabel(xlabel) @@ -547,6 +668,48 @@ 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 twinx(self) -> _Axis: + """Like `Axes.twinx`: this panel's right-hand y-axis, whether or not it holds anything yet.""" + return _Axis(self._plot, self._index, right=True) + + def plot(self, *metrics) -> _Axis: + """Like `ax.plot("name", data=...)`: put these metrics on the left axis. See `_Axis.plot`.""" + return self.left.plot(*metrics) + + def imshow(self, x, *, rows=None, cols=None, griddim=None, vmin=None, vmax=None, scale_each=False, + channels=None, cmap="gray", padding=0, pad_value=0, max_images=64, max_cols=8): + """ + Like `Axes.imshow`, but for a whole batch and repeatable: show `x` on this panel, replacing + whatever was there. A batch is tiled into a grid -- `rows` or `cols` alone infers the other, + `griddim=(rows, cols)` fixes both (padding with blanks, or dropping the tail). With neither, + the grid is near-square (3 -> 2x2, 5 -> 3x2, 8 -> 3x3, 64 -> 8x8), at most `max_cols` wide. + Only the first `max_images` are shown, with a warning (`max_images=None` for all of them). + Images sit edge to edge; `padding` puts that many pixels of `pad_value` between them. The + title shows the step the image was drawn at. + + Accepts (H, W), (C, H, W), (H, W, C), (B, H, W), (B, C, H, W) and (B, H, W, C); (3, H, W) + and (4, H, W) are ambiguous and raise, telling you which `channels=` to pass. + + Values are scaled to the full range by default, over the whole batch; `vmin` / `vmax` fix + the range instead (worth doing for a live view -- otherwise the black point moves every + frame), and `scale_each=True` scales each image on its own, as make_grid does. Unlike + matplotlib, `vmin` / `vmax` are honoured for colour images too. + """ + from ._images import to_grid + + spec = self.spec + assert not (spec["metrics"] or spec["secondary"]), \ + f"panel {self._index} is showing curves ({' '.join(spec['metrics'] + spec['secondary'])}); " \ + f"use a different panel for the image" + arr = to_grid(x, rows=rows, cols=cols, griddim=griddim, vmin=vmin, vmax=vmax, scale_each=scale_each, + channels=channels, padding=padding, pad_value=pad_value, max_images=max_images, + max_cols=max_cols) + if spec["kind"] != "image" or spec["cmap"] != cmap: + spec["kind"], spec["cmap"] = "image", cmap + self._plot._send_layout() + self._plot._send_image(self._index, arr) + return self + def _set(self, key, value): self.spec[key] = value self._plot._send_layout() @@ -554,6 +717,14 @@ def _set(self, key, value): def set_title(self, label): self._set("title", str(label)) + def legend(self, **kwargs): + """ + Like `Axes.legend`, for the panel's one legend (both y-axes' curves and its reference lines): + kwargs go to matplotlib, e.g. `legend(loc="upper left", ncols=3)`, or + `legend(loc="upper center", bbox_to_anchor=(0.5, -0.15), ncols=3)` to put it under the panel. + """ + self._set("legend", kwargs) + def set_xlabel(self, xlabel): self._set("xlabel", str(xlabel)) @@ -592,6 +763,47 @@ def __repr__(self): return f"Panel({self._index}: {' '.join(self.spec['metrics'])}{' | ' + ' '.join(self.spec['secondary']) if self.spec['secondary'] else ''})" +class _PanelGrid: + """ + The panel array `LivePlot.subplots` returns, indexable the way matplotlib's is: `axes[0][1]` + and `axes[0, 1]` both work, `axes.flat` walks it in row-major order, and it unpacks. + """ + + def __init__(self, rows: list): + self._rows = rows + + def __getitem__(self, key): + if isinstance(key, tuple): + row, col = key + return self._rows[row][col] + return self._rows[key] + + def __len__(self): + return len(self._rows) + + def __iter__(self): + return iter(self._rows) + + @property + def flat(self) -> list: + return [panel for row in self._rows for panel in row] + + @property + def shape(self) -> tuple: + return (len(self._rows), len(self._rows[0])) + + def _squeezed(self): + """matplotlib's squeeze: 1x1 -> the panel, 1xN or Nx1 -> a flat list, else the grid.""" + flat = self.flat + if len(flat) == 1: + return flat[0] + rows, cols = self.shape + return flat if rows == 1 or cols == 1 else self + + def __repr__(self): + return f"_PanelGrid({self.shape[0]}x{self.shape[1]})" + + # --------------------------------------------------------------------------- the handle @@ -634,6 +846,7 @@ def __init__( 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 + self._own_bar = None # the bar update() opened, which finish() closes if total is None and iterable is not None: try: total = len(iterable) @@ -651,6 +864,9 @@ def __init__( self._record_path = record if isinstance(record, str) else None self.frames: list[tuple[float, bytes]] = [] # (time, png) of every frame shown, if record=True self.axvlines: list[dict] = [] # vertical reference lines added with axvline() + self._images: dict[int, object] = {} # panel index -> the uint8 array it is showing + self._image_steps: dict[int, object] = {} # panel index -> the x it was shown at + self._fixed_grid = False # set by subplots(): grow a panel's metrics, never the grid self._t0 = time.monotonic() self._done = False self._proc = self._renderer = self._inbox = self._outbox = None @@ -669,7 +885,7 @@ def __init__( f"rendering on the training thread instead (~0.15 s per redraw).", stacklevel=2, ) - self._renderer = _FigureRenderer(self._panels_or_placeholder(), self.x_range, self._layout) + self._renderer = _FigureRenderer(self._panels_or_placeholder(), self.x_range, self._layout, images=self._image_state()) self.mode = "thread" # -- setup ------------------------------------------------------------------- @@ -692,6 +908,43 @@ def _with_defaults(self, panel): panel[key] = value return panel + @classmethod + def subplots(cls, nrows: int = 1, ncols: int = 1, *, figsize=None, squeeze: bool = True, + width_ratios=None, height_ratios=None, iterable=None, **kwargs): + """ + Like `plt.subplots`, for a live plot: returns `(plot, axes)` with a fixed nrows x ncols grid + of empty panels, which you then fill with `ax.plot("loss", ...)` or `ax.imshow(tensor)`. + + plot, (ax_loss, ax_samples) = LivePlot.subplots(1, 2, total=n_steps, figsize=(11, 4)) + ax_loss.plot("lossD", "lossG") + ax_samples.imshow(netG(fixed_noise), rows=2, vmin=-1, vmax=1) + + `axes` follows matplotlib's squeeze rules: one panel for 1x1, a flat list for a single row + or column, a 2-d grid otherwise. `figsize` is the whole figure in inches, as matplotlib + means it (liveplot's own `cell_size` is per panel). `width_ratios` / `height_ratios` are + matplotlib's too, e.g. `width_ratios=(1, 1.3)` to give an image panel more room. Every other + keyword goes to `LivePlot`. + """ + assert nrows >= 1 and ncols >= 1, f"need at least a 1x1 grid, got {nrows}x{ncols}" + kwargs.setdefault("rows", nrows) + kwargs.setdefault("cols", ncols) + if figsize is not None: + kwargs["cell_size"] = (figsize[0] / ncols, figsize[1] / nrows) + plot = cls(*([] if iterable is None else [iterable]), **kwargs) + gridspec_kw = {k: list(v) for k, v in (("width_ratios", width_ratios), ("height_ratios", height_ratios)) if v is not None} + if gridspec_kw: + assert len(gridspec_kw.get("width_ratios", [0] * ncols)) == ncols, "width_ratios needs one entry per column" + assert len(gridspec_kw.get("height_ratios", [0] * nrows)) == nrows, "height_ratios needs one entry per row" + plot._layout = (*plot._layout[:6], gridspec_kw) + plot._specs.extend( + plot._with_defaults(_normalise_panel({"metrics": []}, allow_empty=True)) for _ in range(nrows * ncols) + ) + plot._explicit_layout = True + plot._fixed_grid = True + plot._send_layout() + grid = _PanelGrid([[Panel(plot, r * ncols + c) for c in range(ncols)] for r in range(nrows)]) + return plot, (grid._squeezed() if squeeze else grid) + # -- panels and axes, addressed like matplotlib ----------------------------------- @property @@ -859,6 +1112,32 @@ def __call__(self, iterable, **tqdm_kwargs): if ours: bar.close() + def update(self, n=1): + """ + Like tqdm's `update`, for a loop you drive yourself: advance the count, and so the x-axis, by + `n` items. The first call opens a bar of `total` items below the plot, as `tqdm(total=...)` + would, and `finish()` closes it. One bar for a whole multi-epoch run: + + plot = LivePlot(total=epochs * len(loader)) + for epoch in range(epochs): + for batch in loader: + plot.log(loss=train_step(batch)) + plot.update() + plot.finish() + """ + self._wrapping = True + if self._bar is None and self._own_bar is None: + self._own_bar = self._bar = self._make_bar(None, {"desc": self._desc} if self._desc else {}) + self.n += n + if self._own_bar is not None: + self._own_bar.update(n) + + def set_description(self, desc=None, refresh=True): + """Like tqdm's `set_description`: the text before the bar (e.g. f"epoch {epoch}").""" + self._desc = desc + if self._bar is not None: + self._bar.set_description(desc, refresh=refresh) + def __iter__(self): if self._iterable is None: raise TypeError("nothing to iterate: use LivePlot(iterable, ...), or plot(iterable) inside your loop") @@ -910,20 +1189,32 @@ def _extend_layout(self, new_metrics): unplaced = [m for m in new_metrics if m not in self._placed] if not unplaced: return - if self._explicit_layout: + curves = [i for i, spec in enumerate(self._specs) if spec["kind"] == "curve"] + if self._fixed_grid: + # subplots() promised a grid of this shape, so grow a panel rather than the grid: the + # first one still empty, else the first curve panel there is. + assert curves, "subplots() gave this plot no curve panel to hold " + ", ".join(unplaced) + empty = [i for i in curves if not self._specs[i]["metrics"] and not self._specs[i]["secondary"]] + self._join_panel(empty[0] if empty else curves[0], unplaced) + elif self._explicit_layout: for m in unplaced: self._specs.append(self._with_defaults(_normalise_panel({"metrics": [m]}))) - elif not self._specs: + elif not curves: # nothing but image panels so far (plot.imshow() before the first log) self._specs.append(self._with_defaults(_normalise_panel({"metrics": unplaced}))) else: - spec = self._specs[0] - was_auto = spec["title"] == " / ".join(spec["metrics"] + spec["secondary"]) # a title nobody chose - spec["metrics"].extend(unplaced) - if was_auto: # a title set with set_title() survives a newly discovered metric - spec["title"] = " / ".join(spec["metrics"] + spec["secondary"]) + self._join_panel(curves[0], unplaced) self._placed.update(unplaced) self._send_layout() + def _join_panel(self, index: int, unplaced: list): + """Add metrics to an existing panel, keeping a title the user chose.""" + spec = self._specs[index] + auto = " / ".join(spec["metrics"] + spec["secondary"]) + was_auto = spec["title"] in (auto, "") # a title nobody chose, or a fresh subplots() panel + spec["metrics"].extend(unplaced) + if was_auto: # a title set with set_title() survives a newly discovered metric + spec["title"] = " / ".join(spec["metrics"] + spec["secondary"]) + def axhline(self, y, label=None, *, metric=None, **kwargs): """ Like matplotlib's `Axes.axhline`: a horizontal reference line at `y` with `label` in the legend, @@ -951,10 +1242,44 @@ def axvline(self, x=None, label=None, **kwargs): elif self.mode == "thread": self._renderer.add_axvline(line) + def imshow(self, x, **kwargs): + """ + Show an image on the plot, replacing whatever was there -- the whole of `LiveImage` in one + call. Uses the plot's image panel, creating it on first use. Keywords go to `Panel.imshow`. + """ + for i, spec in enumerate(self._specs): + if spec["kind"] == "image": + return Panel(self, i).imshow(x, **kwargs) + assert not self._fixed_grid, \ + "this plot's grid comes from subplots(); call imshow() on one of its panels instead" + self._specs.append(self._with_defaults(_normalise_panel({"kind": "image"}))) + self._send_layout() + return Panel(self, len(self._specs) - 1).imshow(x, **kwargs) + + def _place(self, metrics): + """Record metrics as already belonging to a panel, and redraw.""" + self._placed.update(metrics) + for name in metrics: + self.data.setdefault(name, ([], [])) + self._send_layout() + + def _send_image(self, index: int, arr): + x = self.step if self._wrapping or self.latest else None # a step in the title only once there are steps + self._images[index], self._image_steps[index] = arr, x + if self.mode == "process": + self._inbox.put(("image", index, arr, x)) + elif self.mode == "thread": + self._renderer.add_image(index, arr, x) + + def _image_state(self) -> dict: + """What a fresh renderer needs to redraw the images: panel index -> (array, x).""" + return {i: (arr, self._image_steps.get(i)) for i, arr in self._images.items()} + def _send_layout(self): if self.mode == "process": - self._inbox.put(("layout", self._specs)) + self._inbox.put(("layout", self._specs, self._layout)) elif self.mode == "thread": + self._renderer.layout = self._layout self._renderer.set_layout(self._specs) def figure(self): @@ -963,7 +1288,7 @@ def figure(self): built on the calling thread and independent of the render process: title it, tweak it, `fig.savefig("run.png")`, or show it in a report. Safe to call during or after training. """ - renderer = _FigureRenderer(self._panels_or_placeholder(), self.x_range, self._layout) + renderer = _FigureRenderer(self._panels_or_placeholder(), self.x_range, self._layout, images=self._image_state()) for name, (xs, ys) in self.data.items(): renderer.hist[name] = (list(xs), list(ys)) for line in self.axvlines: @@ -994,6 +1319,9 @@ def finish(self): if self._proc.is_alive(): # e.g. still importing matplotlib on a very busy machine self._proc.terminate() self._proc.join(timeout=2.0) + if self._own_bar is not None: + self._own_bar.set_postfix(self.latest, refresh=False) + self._own_bar.close() if self._record_path and self.frames: self.save_gif(self._record_path) @@ -1041,7 +1369,7 @@ def _fall_back_to_thread(self, why): # process will usually kill a renderer built here too (a bad panel spec, a missing backend), # so if this fails, say so once and carry on collecting into plot.data. try: - renderer = _FigureRenderer(self._panels_or_placeholder(), self.x_range, self._layout) + renderer = _FigureRenderer(self._panels_or_placeholder(), self.x_range, self._layout, images=self._image_state()) for name, (xs, ys) in self.data.items(): renderer.hist[name] = (list(xs), list(ys)) for line in self.axvlines: diff --git a/tests/test_images.py b/tests/test_images.py new file mode 100644 index 0000000..979c38a --- /dev/null +++ b/tests/test_images.py @@ -0,0 +1,144 @@ +"""Shape reading, grid tiling and value scaling for `imshow` -- all of `_images.py`, no figure.""" + +import warnings + +import numpy as np +import pytest + +from liveplot._images import grid_shape, to_grid + + +def test_accepted_layouts(): + """Every layout, read into one picture. (H, W) and (H, W, 3) come out with no batch dimension.""" + assert to_grid(np.zeros((28, 28))).shape == (28, 28) + assert to_grid(np.zeros((1, 28, 28))).shape == (28, 28), "(1, H, W): both readings draw this" + assert to_grid(np.zeros((64, 64, 3))).shape == (64, 64, 3), "channels last" + assert to_grid(np.zeros((64, 64, 4))).shape == (64, 64, 4) + assert to_grid(np.zeros((3, 64, 64)), channels="first").shape == (64, 64, 3) + assert to_grid(np.zeros((5, 8, 8))).shape == (24, 16), "5 grayscale -> near-square 3x2, one blank" + assert to_grid(np.zeros((4, 1, 8, 8))).shape == (16, 16), "(B, 1, H, W) -> 2x2, edge to edge" + assert to_grid(np.zeros((4, 3, 8, 8))).shape == (16, 16, 3), "(B, C, H, W), torch's layout" + assert to_grid(np.zeros((4, 8, 8, 3))).shape == (16, 16, 3), "(B, H, W, C), channels last" + + +def test_the_ambiguous_shapes_raise_and_say_how_to_fix_it(): + """matplotlib dodges this by refusing channels-first entirely; we have to accept it, so we ask.""" + for n, kind in ((3, "RGB"), (4, "RGBA")): + with pytest.raises(ValueError, match="ambiguous shape") as e: + to_grid(np.zeros((n, 8, 8))) + assert kind in str(e.value) and 'channels="first"' in str(e.value) and 'channels="none"' in str(e.value) + # ... and both escape hatches work, giving genuinely different pictures + assert to_grid(np.zeros((3, 8, 8)), channels="first").shape == (8, 8, 3) + assert to_grid(np.zeros((3, 8, 8)), channels="none", padding=0).shape == (16, 16), "3 grayscale -> 2x2, one blank" + with pytest.raises(ValueError, match="2-, 3- or 4-d"): + to_grid(np.zeros((2, 3, 4, 5, 6))) + with pytest.raises(AssertionError, match="not a batch of images"): + to_grid(np.zeros((4, 7, 8, 9))) # no dimension that could be channels + + +def test_grid_shape_inference(): + assert grid_shape(10, rows=2, cols=None) == (2, 5) + assert grid_shape(10, rows=None, cols=5) == (2, 5) + assert grid_shape(10, rows=None, cols=None) == (4, 3), "near-square, as for the panel grid" + assert grid_shape(10, rows=None, cols=None, griddim=(2, 5)) == (2, 5) + assert grid_shape(7, rows=3, cols=None) == (3, 3), "rounds up, leaving blanks" + with pytest.raises(AssertionError, match="not both"): + grid_shape(4, rows=2, cols=2, griddim=(2, 2)) + + +def test_grid_pads_short_and_drops_long(): + imgs = np.arange(3 * 4 * 4, dtype=np.float32).reshape(3, 4, 4) # 3 distinct grayscale images + padded = to_grid(imgs, griddim=(2, 2), channels="none", padding=0) + assert padded.shape == (8, 8) + assert (padded[4:, 4:] == 0).all(), "the 4th cell is a blank, not a repeat" + dropped = to_grid(imgs, griddim=(1, 2), channels="none", padding=0) + assert dropped.shape == (4, 8) + assert np.array_equal(dropped, to_grid(imgs[:2], griddim=(1, 2), channels="none", padding=0)), \ + "the tail is dropped before scaling, so an image nobody sees cannot set the range" + + +def test_tiling_puts_images_in_row_major_order(): + imgs = np.stack([np.full((2, 2), v, np.float32) for v in (0.0, 1.0, 2.0, 3.0)]) + g = to_grid(imgs, rows=2, cols=2, channels="none", vmin=0, vmax=3, padding=0) + assert [g[0, 0], g[0, 3], g[3, 0], g[3, 3]] == [0, 85, 170, 255], "0 1 / 2 3, reading across" + + +def test_value_scaling(): + x = np.array([[-1.0, 0.0, 1.0]], dtype=np.float32) + assert to_grid(x).tolist() == [[0, 128, 255]], "no range given: min-max over the batch" + assert to_grid(x, vmin=-1, vmax=1).tolist() == [[0, 128, 255]], "the same range, fixed" + assert to_grid(x, vmin=0, vmax=1).tolist() == [[0, 0, 255]], "out-of-range values clip" + flat = np.full((1, 4, 4), 0.7, np.float32) + assert (to_grid(flat) == 0).all(), "a constant image comes out at 0, as in make_grid" + + u8 = np.array([[0, 7, 255]], dtype=np.uint8) + assert to_grid(u8).tolist() == [[0, 7, 255]], "uint8 passes through untouched" + assert to_grid(u8, vmin=0, vmax=7).tolist() == [[0, 255, 255]], "... unless a range is given" + + +def test_scale_each(): + """Two images with very different ranges: together the dim one vanishes, apart both fill 0..255.""" + imgs = np.stack([np.array([[0.0, 0.1]]), np.array([[0.0, 1.0]])]).astype(np.float32) + together = to_grid(imgs, rows=1, cols=2, channels="none", padding=0) + apart = to_grid(imgs, rows=1, cols=2, channels="none", scale_each=True, padding=0) + assert together.tolist() == [[0, 26, 0, 255]], "one global range" + assert apart.tolist() == [[0, 255, 0, 255]], "each image on its own range" + assert to_grid(imgs, rows=1, cols=2, channels="none", scale_each=True, vmin=0, vmax=1, padding=0).tolist() \ + == [[0, 26, 0, 255]], "an explicit range wins over scale_each" + + +def test_torch_tensors_without_importing_torch(): + """`to_grid` duck-types .detach()/.cpu(); nothing in liveplot names torch.""" + class FakeTensor: + def __init__(self, a): self.a, self.detached, self.moved = a, False, False + def detach(self): self.detached = True; return self + def cpu(self): self.moved = True; return self + def __array__(self, dtype=None, copy=None): return self.a + + x = FakeTensor(np.zeros((4, 1, 8, 8), np.float32)) + assert to_grid(x, padding=0).shape == (16, 16) + assert x.detached and x.moved, "a grad-tracking GPU tensor must be brought back first" + import liveplot.liveplot as lp + assert "torch" not in str(lp.__dict__.keys()) + + +def test_pad_value(): + imgs = np.ones((1, 4, 4), np.float32) + assert (to_grid(imgs, griddim=(1, 2), channels="none", pad_value=255, padding=0)[:, 4:] == 255).all() + + +def test_padding_between_images(): + """make_grid's layout: `padding` pixels around every image; one image alone gets no border.""" + g = to_grid(np.ones((2, 1, 3, 3), np.float32), griddim=(1, 2), vmin=0, vmax=1, padding=2, pad_value=7) + assert g.shape == (3 + 4, 2 * 3 + 6) + assert (g[2:5, 2:5] == 255).all() and (g[2:5, 7:10] == 255).all(), "the images" + assert (g[:2] == 7).all() and (g[:, 5:7] == 7).all(), "the border and the gap between" + assert to_grid(np.ones((3, 3)), padding=2).shape == (3, 3) + assert to_grid(np.ones((2, 1, 3, 3)), griddim=(1, 2)).shape == (3, 6), "the default is edge to edge" + + +def test_default_grid_shapes(): + """Near-square, rows first: what "just plot this batch" gives, capped at 8 columns.""" + shapes = {n: grid_shape(n, None, None) for n in (1, 2, 3, 4, 5, 8, 9, 10, 16, 64)} + assert shapes == {1: (1, 1), 2: (2, 1), 3: (2, 2), 4: (2, 2), 5: (3, 2), 8: (3, 3), 9: (3, 3), + 10: (4, 3), 16: (4, 4), 64: (8, 8)} + assert grid_shape(100, None, None) == (13, 8), "never wider than max_cols=8 ..." + assert grid_shape(100, None, None, max_cols=None) == (10, 10), "... unless told otherwise" + assert grid_shape(100, None, 20) == (5, 20), "an explicit cols is not capped" + + +def test_max_images_truncates_with_a_warning(): + many = np.zeros((1000, 1, 2, 2), np.float32) + with pytest.warns(UserWarning, match="first 64 of 1000"): + assert to_grid(many).shape == (16, 16), "64 images, 8x8" + with warnings.catch_warnings(): + warnings.simplefilter("error") + assert to_grid(many, max_images=None).shape == (125 * 2, 8 * 2), "all 1000, 8 wide: asked for, so no warning" + assert to_grid(many[:64]).shape == (16, 16), "exactly max_images: no warning" + + +def test_bfloat16_tensors(): + """What autocast hands back; numpy has no bfloat16, so it must be widened before np.asarray.""" + torch = pytest.importorskip("torch") + x = torch.rand(4, 3, 5, 5, dtype=torch.bfloat16, requires_grad=True) + assert np.array_equal(to_grid(x), to_grid(x.detach().float().numpy())) diff --git a/tests/test_subplots.py b/tests/test_subplots.py new file mode 100644 index 0000000..a0fb4e8 --- /dev/null +++ b/tests/test_subplots.py @@ -0,0 +1,262 @@ +"""`LivePlot.subplots`, `ax.plot`/`ax.twinx`, and image panels sharing a figure with curves.""" + +import time + +import numpy as np +import pytest + +from liveplot import LivePlot +from liveplot.liveplot import Panel, _FigureRenderer, _normalise_panel, _PanelGrid + +LAYOUT = (3, None, None, (4, 3), 50, "step") +PNG = b"\x89PNG\r\n\x1a\n" + + +def rng_images(n=4, c=1, h=8, w=8, seed=0): + return np.random.default_rng(seed).random((n, c, h, w)).astype(np.float32) + + +def test_subplots_squeezes_like_matplotlib(): + plot, ax = LivePlot.subplots(progress=False) + assert isinstance(ax, Panel), "1x1 -> the panel itself" + plot, axs = LivePlot.subplots(1, 3, progress=False) + assert isinstance(axs, list) and len(axs) == 3, "a single row -> a flat list, so it unpacks" + plot, axs = LivePlot.subplots(3, 1, progress=False) + assert isinstance(axs, list) and len(axs) == 3 + plot, axs = LivePlot.subplots(2, 2, progress=False) + assert isinstance(axs, _PanelGrid) and axs.shape == (2, 2) and len(axs.flat) == 4 + assert axs[0, 1].spec is axs[0][1].spec, "matplotlib's two spellings agree" + assert axs[1, 0].spec is plot.panels[2].spec, "row-major, as in matplotlib" + plot, grid = LivePlot.subplots(1, 2, squeeze=False, progress=False) + assert isinstance(grid, _PanelGrid) and grid.shape == (1, 2) + with pytest.raises(AssertionError, match="1x1 grid"): + LivePlot.subplots(0, 2, progress=False) + + +def test_figsize_is_the_whole_figure_as_matplotlib_means_it(): + plot, _ = LivePlot.subplots(2, 4, figsize=(12, 6), progress=False) + assert plot._layout[3] == (3.0, 3.0), "cell_size is figsize / grid" + assert plot._layout[1:3] == (2, 4), "the grid is fixed at what was asked for" + + +def test_plot_assigns_metrics_and_twinx_gives_the_right_axis(): + plot, (ax_loss, ax_lr) = LivePlot.subplots(1, 2, progress=False) + returned = ax_loss.plot("lossD", "lossG") + assert returned.plot("lossD") is returned, "chainable, and a repeat is not a duplicate" + ax_loss.twinx().plot("D(x)") + ax_lr.plot("lr") + assert ax_loss.spec["metrics"] == ["lossD", "lossG"] and ax_loss.spec["secondary"] == ["D(x)"] + assert ax_loss.spec["title"] == "lossD / lossG / D(x)", "the auto title follows both axes" + plot.log(0, lossD=1.0, lossG=0.5, lr=1e-3, **{"D(x)": 0.6}) + assert len(plot._specs) == 2, "subplots() fixed the grid: no third panel appeared" + assert plot["D(x)"]._right and plot["lossD"]._right is False + with pytest.raises(AssertionError, match="metric names, not data"): + ax_lr.plot([1, 2, 3]) + + +def test_a_metric_nobody_declared_joins_a_panel_rather_than_growing_the_grid(): + plot, (ax_loss, ax_img) = LivePlot.subplots(1, 2, progress=False) + ax_loss.plot("lossD") + ax_img.imshow(rng_images()) + plot.log(0, lossD=1.0, surprise=2.0) + assert len(plot._specs) == 2, "the grid subplots() promised is not resized" + assert plot._specs[0]["metrics"] == ["lossD", "surprise"], "it lands on the curve panel" + # with an empty curve panel available, that one is preferred + plot2, (a, b) = LivePlot.subplots(1, 2, progress=False) + a.plot("loss") + plot2.log(0, loss=1.0, other=2.0) + assert (plot2._specs[0]["metrics"], plot2._specs[1]["metrics"]) == (["loss"], ["other"]) + # a grid of nothing but image panels has nowhere to put it + plot3, ax = LivePlot.subplots(progress=False) + ax.imshow(rng_images()) + with pytest.raises(AssertionError, match="no curve panel"): + plot3.log(0, loss=1.0) + + +def test_image_panel_and_curve_panel_do_not_mix(): + plot, (ax_loss, ax_img) = LivePlot.subplots(1, 2, progress=False) + ax_loss.plot("loss") + with pytest.raises(AssertionError, match="showing curves"): + ax_loss.imshow(rng_images()) + ax_img.imshow(rng_images()) + with pytest.raises(AssertionError, match="showing an image"): + ax_img.plot("loss") + + +def test_plot_level_imshow_is_the_whole_of_liveimage(): + plot = LivePlot(progress=False) + plot.imshow(rng_images(seed=1)) + assert plot._image_steps[0] is None, "nothing logged or wrapped yet: no step to put in the title" + first = plot._images[0].copy() + plot.imshow(rng_images(seed=2)) + assert [s["kind"] for s in plot._specs] == ["image"], "one image panel, reused" + assert not np.array_equal(first, plot._images[0]), "the second call overwrites the first" + plot.log(0, loss=1.0) # a curve discovered after the image panel exists + assert [s["kind"] for s in plot._specs] == ["image", "curve"] + assert plot._specs[1]["metrics"] == ["loss"] + q, ax = LivePlot.subplots(progress=False) + with pytest.raises(AssertionError, match="comes from subplots"): + q.imshow(rng_images()) + + +def test_renderer_draws_images_and_keeps_them_through_a_relayout(): + panels = [_normalise_panel({"metrics": ["loss"]}), _normalise_panel({"kind": "image", "title": "samples"})] + r = _FigureRenderer(panels, xlim=None, layout=LAYOUT) + for step in range(5): + r.add(step, {"loss": 1.0 / (step + 1)}) + arr = (np.arange(16 * 16, dtype=np.uint8).reshape(16, 16)) + r.add_image(1, arr) + assert r.render()[:8] == PNG + ax_img = r.image_axes[1] + assert len(ax_img.images) == 1 and ax_img.get_title() == "samples" + assert not ax_img.axison, "pixel indices along a tiled grid are noise" + artist = r.image_artists[1] + r.add_image(1, arr[::-1].copy()) + assert r.image_artists[1] is artist, "same shape: the artist is reused, not rebuilt" + r.add_image(1, np.zeros((8, 8), np.uint8)) + assert r.image_artists[1] is not artist, "a new shape needs a new artist" + # a re-layout rebuilds the figure; the picture must come back with it + r.set_layout(panels + [_normalise_panel("lr")]) + assert len(r.image_axes[1].images) == 1 and r.render()[:8] == PNG + assert len(r.hist["loss"][0]) == 5, "and the curve history survives as before" + + +def test_axvline_skips_image_panels(): + panels = [_normalise_panel({"metrics": ["loss"]}), _normalise_panel({"kind": "image"})] + r = _FigureRenderer(panels, xlim=None, layout=LAYOUT) + r.add_image(1, np.zeros((8, 8), np.uint8)) + r.add_axvline({"x": 2, "label": "lr drop"}) + assert any(line.get_linestyle() == ":" for line in r.axes[0].get_lines()) + assert not r.image_axes[1].get_lines(), "a vertical line across a picture means nothing" + + +def test_figure_includes_the_image(): + plot, (ax_loss, ax_img) = LivePlot.subplots(1, 2, total=20, progress=False) + ax_loss.plot("loss") + for step in range(5): + plot.log(step, loss=1.0 / (step + 1)) + ax_img.imshow(rng_images(n=4, c=3), rows=2, vmin=0, vmax=1) + fig = plot.figure() + assert sum(len(ax.images) for ax in fig.axes) == 1 + assert sum(len(ax.get_lines()) for ax in fig.axes) == 1 + img = next(im for ax in fig.axes for im in ax.images) + assert img.get_array().shape == (16, 16, 3), "two rows of two 8x8 RGB images" + assert img.axes.get_title() == "step 4", "an untitled image panel says when it was drawn" + ax_img.set_title("samples") + assert next(im for ax in plot.figure().axes for im in ax.images).axes.get_title() == "samples (step 4)" + + +class _FakeHandle: + def __init__(self): + self.frames = [] + + def update(self, img): + self.frames.append(img.data) + + +@pytest.fixture +def fake_notebook(monkeypatch): + h = _FakeHandle() + monkeypatch.setattr(LivePlot, "_make_display_handle", staticmethod(lambda: h)) + return h + + +def _colour_images_are_safe_here() -> bool: + """ + matplotlib 3.9.0 corrupts its multi-channel resample path when the figure is drawn from a + spawned render process: any RGB(A) image raises "arrays must be of dtype byte, short, float32 + or float64" from the C extension, while the same array in the same figure draws fine in the + parent, and grayscale is unaffected. The failure flips on perturbations as transparent as + wrapping `matplotlib.image._resample`, which is what memory corruption looks like rather than + a bug of ours. Fixed by 3.9.4; reproducible with numpy 1.26 and 2.4 alike, so it is matplotlib, + not the numpy pairing. Colour images are still covered in-process by test_figure_includes_the_image. + """ + from packaging.version import Version + + import matplotlib + + return not (Version("3.9.0") <= Version(matplotlib.__version__) < Version("3.9.4")) + + +@pytest.mark.skipif(not _colour_images_are_safe_here(), reason="matplotlib 3.9.0-3.9.3 breaks RGB draws in a spawned process") +def test_process_mode_curves_beside_images(fake_notebook): + """The DCGAN shape: curves every step, samples every so often, one figure, one output cell.""" + plot, (ax_loss, ax_img) = LivePlot.subplots(1, 2, total=200, refresh_seconds=0.1, figsize=(8, 3), dpi=40) + ax_loss.plot("lossD", "lossG") + ax_img.set_title("generator samples") + assert plot.mode == "process" + t0 = time.monotonic() + while not fake_notebook.frames and time.monotonic() - t0 < 90: # child startup + plot.log(0, lossD=1.0, lossG=1.0) + time.sleep(0.05) + assert fake_notebook.frames, "render child produced no frame" + n = len(fake_notebook.frames) + for step in range(60): + plot.log(step, lossD=1.0 / (step + 1), lossG=0.5) + if step % 20 == 0: # tanh-range samples, as a generator produces + ax_img.imshow(rng_images(n=4, c=3, seed=step) * 2 - 1, rows=2, vmin=-1, vmax=1) + time.sleep(0.02) + plot.finish() + assert len(fake_notebook.frames) > n and all(f[:8] == PNG for f in fake_notebook.frames) + assert not plot._proc.is_alive() + assert len(plot._images) == 1 and plot._images[1].shape == (16, 16, 3) + assert [s["kind"] for s in plot._specs] == ["curve", "image"] + + +def test_width_ratios_and_legend_placement(): + """matplotlib's own names: subplots(width_ratios=...) and ax.legend(**kwargs).""" + plot, (ax_loss, ax_img) = LivePlot.subplots(1, 2, width_ratios=(1, 2), progress=False) + ax_loss.plot("lossD", "lossG") + ax_loss.twinx().plot("D(x)") + ax_loss.legend(loc="upper center", ncols=3) + ax_img.imshow(rng_images()) + plot.log(0, lossD=1.0, lossG=0.5, **{"D(x)": 0.5}) + fig = plot.figure() + axes = [ax for ax in fig.axes if ax.get_visible()] + ax_curves, ax_image = axes[0], axes[1] + assert ax_image.get_subplotspec().get_gridspec().get_width_ratios() == [1, 2] + legend = ax_curves.get_legend() + assert [t.get_text() for t in legend.get_texts()] == ["lossD", "lossG", "D(x)"], "one legend, both axes' curves" + assert legend._ncols == 3 and legend._loc == 9, "kwargs reached Axes.legend (9 = upper center)" + with pytest.raises(AssertionError, match="one entry per column"): + LivePlot.subplots(1, 2, width_ratios=(1, 2, 3), progress=False) + + +def test_update_drives_one_bar_across_epochs(): + """tqdm's manual mode: plot.update() advances the count (the x-axis) and one bar for the run.""" + pytest.importorskip("tqdm") + plot = LivePlot(total=6) + bars = set() + for epoch in range(2): + plot.set_description(f"epoch {epoch}") + for _ in range(3): + plot.log(loss=1.0) + plot.update() + bars.add(id(plot._bar)) + plot.finish() + assert len(bars) == 1 and plot._own_bar.n == 6 and plot._own_bar.total == 6 + assert plot._own_bar.desc.startswith("epoch 1") + assert plot.data["loss"][0] == [0, 1, 2, 3, 4, 5], "x counts updates, like a wrapped loop" + quiet = LivePlot(total=3, progress=False) + quiet.update(); quiet.update(2) + assert quiet._bar is None and quiet.step == 3, "progress=False: the count moves, no bar" + + +def test_image_panels_are_sized_to_their_picture(): + """One row, no width_ratios: the image panel is as wide as its picture needs, the curves get the rest.""" + plot, (ax_loss, ax_img) = LivePlot.subplots(1, 2, figsize=(12, 4), progress=False) + ax_loss.plot("loss") + plot.log(0, loss=1.0) + ax_img.imshow(rng_images(n=9, c=3)) # 3x3 of 8x8: square + fig = plot.figure() + assert fig.get_size_inches()[0] == pytest.approx(12), "figsize is kept" + ax_curves, ax_image = [ax for ax in fig.axes if ax.get_visible()] + fig.canvas.draw() + box = ax_image.get_window_extent() + assert box.width == pytest.approx(box.height, rel=0.02), "the picture fills its panel: no side margins" + assert ax_curves.get_window_extent().width > 1.5 * box.width, "the curves take the room the picture doesn't need" + ax_img.imshow(rng_images(n=8, c=3), griddim=(2, 4)) # a wide picture: the panel re-fits + box = [ax for ax in plot.figure().axes if ax.get_visible()][1] + box.figure.canvas.draw() + extent = box.get_window_extent() + assert extent.width == pytest.approx(2 * extent.height, rel=0.02)