diff --git a/examples/eval-tables/.gitignore b/examples/eval-tables/.gitignore new file mode 100644 index 00000000..689af916 --- /dev/null +++ b/examples/eval-tables/.gitignore @@ -0,0 +1 @@ +MNIST/ diff --git a/examples/eval-tables/README.md b/examples/eval-tables/README.md new file mode 100644 index 00000000..20cf5d0a --- /dev/null +++ b/examples/eval-tables/README.md @@ -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. diff --git a/examples/eval-tables/eval_tables_demo.py b/examples/eval-tables/eval_tables_demo.py new file mode 100644 index 00000000..5a7579b3 --- /dev/null +++ b/examples/eval-tables/eval_tables_demo.py @@ -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() diff --git a/examples/eval-tables/mnist.py b/examples/eval-tables/mnist.py new file mode 100644 index 00000000..e83cadc2 --- /dev/null +++ b/examples/eval-tables/mnist.py @@ -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 diff --git a/examples/eval-tables/requirements.txt b/examples/eval-tables/requirements.txt new file mode 100644 index 00000000..10cd0d57 --- /dev/null +++ b/examples/eval-tables/requirements.txt @@ -0,0 +1,3 @@ +torch>=2.1 +torchvision>=0.16 +wandb