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
1 change: 1 addition & 0 deletions examples/eval-tables/.gitignore
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
MNIST/
32 changes: 32 additions & 0 deletions examples/eval-tables/README.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,32 @@
# Eval Tables

> [!NOTE]
> Eval Tables is currently in **Public Preview**. Functionality and availability
> are subject to change.

For an overview of the feature, see the
[Eval Tables Documentation](https://docs.wandb.ai/models/evaltables).

`eval_tables_demo.py` trains a small MNIST classifier with three configs, one
W&B run per config, logging a `wandb.EvalTable` of predictions at the end of
each epoch. In the UI, you will be able to compare runs and specific steps
within each run.

## Setup

```sh
uv venv .venv --python 3.12
source .venv/bin/activate
uv pip install -r requirements.txt
```

## Run

```sh
wandb login
python eval_tables_demo.py
```

The script logs to the `eval-tables-demo` project by default.. Pass `--entity`,
`--project`, `--epochs` or `--val-rows` to change the defaults, and see
`--help` for the rest. MNIST downloads to `./MNIST` on the first run.
105 changes: 105 additions & 0 deletions examples/eval-tables/eval_tables_demo.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,105 @@
"""Log W&B EvalTables across several runs and epochs.

Trains a small MNIST classifier with a few configs. After each epoch, each run
logs an EvalTable of predictions on the same validation images, so the Eval
Tables panel can match rows across steps and runs.
"""

import argparse

import torch
import wandb
from torch.utils.data import DataLoader

from mnist import build_model, evaluate, load_data, train_epoch

INPUT_COLUMNS = ["image", "label"]
OUTPUT_COLUMNS = ["pred", "confidence"]
SCORE_COLUMNS = ["correct", "loss"]

RUN_CONFIGS = [
{"name": "small", "hidden_size": 16, "lr": 1e-3},
{"name": "wide", "hidden_size": 128, "lr": 1e-3},
{"name": "wide-fast", "hidden_size": 128, "lr": 1e-2},
]


def to_uint8(image):
# Identical pixels give identical media hashes, which rows match on.
return (image[0] * 255).round().to(torch.uint8).numpy()


def build_eval_table(images, labels, preds, confidence, losses):
rows = [
[wandb.Image(to_uint8(image)), label, pred, conf, pred == label, loss]
for image, label, pred, conf, loss in zip(
images,
labels.tolist(),
preds.tolist(),
confidence.tolist(),
losses.tolist(),
)
]
return wandb.EvalTable(
columns=[*INPUT_COLUMNS, *OUTPUT_COLUMNS, *SCORE_COLUMNS],
data=rows,
input_columns=INPUT_COLUMNS,
output_columns=OUTPUT_COLUMNS,
score_columns=SCORE_COLUMNS,
)


def run_config(config, args, train_ds, val_images, val_labels):
torch.manual_seed(args.seed)
model = build_model(config["hidden_size"])
optimizer = torch.optim.Adam(model.parameters(), lr=config["lr"])
loader = DataLoader(
train_ds,
batch_size=args.batch_size,
shuffle=True,
generator=torch.Generator().manual_seed(args.seed),
)

with wandb.init(
project=args.project,
entity=args.entity,
name=config["name"],
config={**config, "epochs": args.epochs, "val_rows": args.val_rows},
) as run:
for epoch in range(1, args.epochs + 1):
train_loss = train_epoch(model, loader, optimizer)
preds, confidence, losses = evaluate(model, val_images, val_labels)
run.log(
{
"epoch": epoch,
"train_loss": train_loss,
"val_loss": losses.mean().item(),
"val_accuracy": (preds == val_labels).float().mean().item(),
"val_predictions": build_eval_table(
val_images, val_labels, preds, confidence, losses
),
}
)


def main():
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--project", default="eval-tables-demo")
parser.add_argument("--entity", default=None)
parser.add_argument("--epochs", type=int, default=3)
parser.add_argument("--val-rows", type=int, default=64)
parser.add_argument("--train-size", type=int, default=6000)
parser.add_argument("--batch-size", type=int, default=64)
parser.add_argument("--seed", type=int, default=0)
parser.add_argument("--data-dir", default=".")
args = parser.parse_args()

train_ds, val_images, val_labels = load_data(
args.data_dir, args.train_size, args.val_rows
)
for config in RUN_CONFIGS:
run_config(config, args, train_ds, val_images, val_labels)


if __name__ == "__main__":
main()
50 changes: 50 additions & 0 deletions examples/eval-tables/mnist.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,50 @@
"""MNIST data, model and training helpers for eval_tables_demo.py."""

import torch
import torch.nn.functional as F
from torch import nn
from torch.utils.data import Subset
from torchvision import datasets, transforms

NUM_CLASSES = 10


def load_data(data_dir, train_size, val_rows):
"""Return a training subset and a fixed validation batch."""
to_tensor = transforms.ToTensor()
train = datasets.MNIST(data_dir, train=True, download=True, transform=to_tensor)
val = datasets.MNIST(data_dir, train=False, download=True, transform=to_tensor)
val_images = torch.stack([val[i][0] for i in range(val_rows)])
val_labels = torch.tensor([val[i][1] for i in range(val_rows)])
return Subset(train, range(train_size)), val_images, val_labels


def build_model(hidden_size):
return nn.Sequential(
nn.Flatten(),
nn.Linear(28 * 28, hidden_size),
nn.ReLU(),
nn.Linear(hidden_size, NUM_CLASSES),
)


def train_epoch(model, loader, optimizer):
model.train()
total_loss = 0.0
for images, labels in loader:
optimizer.zero_grad()
loss = F.cross_entropy(model(images), labels)
loss.backward()
optimizer.step()
total_loss += loss.item() * labels.size(0)
return total_loss / len(loader.dataset)


def evaluate(model, images, labels):
"""Return per-row predictions, confidences and losses."""
model.eval()
with torch.inference_mode():
logits = model(images)
losses = F.cross_entropy(logits, labels, reduction="none")
confidence, preds = logits.softmax(dim=1).max(dim=1)
return preds, confidence, losses
3 changes: 3 additions & 0 deletions examples/eval-tables/requirements.txt
Original file line number Diff line number Diff line change
@@ -0,0 +1,3 @@
torch>=2.1
torchvision>=0.16
wandb
Loading