Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
90 changes: 90 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down
Binary file added docs/imshow.gif
Loading
Sorry, something went wrong. Reload?
Sorry, we cannot display this file.
Sorry, this file is invalid so it cannot be displayed.
164 changes: 164 additions & 0 deletions examples/dcgan_synthetic.py
Original file line number Diff line number Diff line change
@@ -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}")
56 changes: 56 additions & 0 deletions examples/demo.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Loading
Loading