Tune a CIFAR-10 model with Ray Tune

Use Ray Tune and the asynchronous successive halving algorithm (ASHA) to search for hyperparameters for a PyTorch image classifier on attached 1xA10 AI Runtime compute. This notebook shows how to:

  • Start Ray on attached GPU compute.
  • Define a checkpointed CIFAR-10 training function.
  • Run four Ray Tune trials, with two fractional-GPU trials running concurrently.
  • Select and evaluate the best trial.

Note

This example requires the Databricks AI environment version 6 or above.

Prerequisites

Connect this notebook to AI Runtime GPU compute:

  1. Select Connect at the top of the notebook.
  2. Select Serverless GPU.
  3. In the Environment side panel, set Accelerator to 1xA10.
  4. Select AI v6 as the base environment.
  5. Select Apply, then select Confirm.

AI v6 includes Ray, PyTorch, and torchvision, so this example does not install additional packages. The first run downloads the CIFAR-10 dataset.

Initialize Ray

Start Ray for the notebook session with ray_init(). The function prints a dashboard link that works through the Databricks driver proxy.

import ray
from serverless_gpu import ray_init

ray_init()

cluster_resources = ray.cluster_resources()
if cluster_resources.get("GPU", 0) < 1:
    raise RuntimeError(
        "Ray did not detect a GPU. Attach the notebook to 1xA10 compute, then rerun it."
    )
cluster_resources

The example limits each trial to a subset of CIFAR-10 and five training epochs to reduce the search runtime. Each trial reserves half of the A10 GPU, which lets Ray schedule up to two trials concurrently.

import tempfile
import uuid
from pathlib import Path

import torch
import torch.nn as nn
import torch.nn.functional as F
from filelock import FileLock
from ray import tune
from ray.tune import Checkpoint
from ray.tune.schedulers import ASHAScheduler
from torch.utils.data import DataLoader, random_split
from torchvision import datasets, transforms

SEED = 42
MAX_EPOCHS = 5
NUM_SAMPLES = 4
CPUS_PER_TRIAL = 2
GPUS_PER_TRIAL = 0.5
TRAIN_SAMPLE_COUNT = 10_000
VALIDATION_SAMPLE_COUNT = 2_000
TEST_SAMPLE_COUNT = 2_000

DATA_DIR = Path(tempfile.gettempdir()) / "ray-tune-cifar10-data"
RESULTS_DIR = Path(tempfile.gettempdir()) / "ray-tune-cifar10-results"

Prepare the dataset and model

The file lock prevents concurrent trials from downloading CIFAR-10 into the same directory at the same time. Every trial uses the same deterministic training and validation subsets.

def load_cifar10(data_dir: Path):
    transform = transforms.Compose(
        [
            transforms.ToTensor(),
            transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5)),
        ]
    )

    data_dir.mkdir(parents=True, exist_ok=True)
    with FileLock(str(data_dir / "download.lock")):
        full_train_dataset = datasets.CIFAR10(
            root=data_dir, train=True, download=True, transform=transform
        )
        full_test_dataset = datasets.CIFAR10(
            root=data_dir, train=False, download=True, transform=transform
        )

    split_generator = torch.Generator().manual_seed(SEED)
    remaining_train_count = (
        len(full_train_dataset) - TRAIN_SAMPLE_COUNT - VALIDATION_SAMPLE_COUNT
    )
    train_dataset, validation_dataset, _ = random_split(
        full_train_dataset,
        [TRAIN_SAMPLE_COUNT, VALIDATION_SAMPLE_COUNT, remaining_train_count],
        generator=split_generator,
    )
    test_dataset, _ = random_split(
        full_test_dataset,
        [TEST_SAMPLE_COUNT, len(full_test_dataset) - TEST_SAMPLE_COUNT],
        generator=torch.Generator().manual_seed(SEED),
    )
    return train_dataset, validation_dataset, test_dataset


class CifarNet(nn.Module):
    def __init__(self, first_hidden_size: int, second_hidden_size: int):
        super().__init__()
        self.conv1 = nn.Conv2d(3, 6, 5)
        self.pool = nn.MaxPool2d(2, 2)
        self.conv2 = nn.Conv2d(6, 16, 5)
        self.fc1 = nn.Linear(16 * 5 * 5, first_hidden_size)
        self.fc2 = nn.Linear(first_hidden_size, second_hidden_size)
        self.fc3 = nn.Linear(second_hidden_size, 10)

    def forward(self, inputs):
        inputs = self.pool(F.relu(self.conv1(inputs)))
        inputs = self.pool(F.relu(self.conv2(inputs)))
        inputs = torch.flatten(inputs, 1)
        inputs = F.relu(self.fc1(inputs))
        inputs = F.relu(self.fc2(inputs))
        return self.fc3(inputs)

Define the Ray Tune training function

Ray Tune calls this function for each sampled hyperparameter configuration. At the end of every epoch, the function reports validation metrics and a checkpoint. ASHA uses the reported loss to stop underperforming trials early.

def train_cifar(config, data_dir: Path, max_epochs: int):
    torch.manual_seed(SEED)
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")

    model = CifarNet(config["first_hidden_size"], config["second_hidden_size"])
    model.to(device)
    criterion = nn.CrossEntropyLoss()
    optimizer = torch.optim.SGD(
        model.parameters(), lr=config["learning_rate"], momentum=0.9
    )
    start_epoch = 0

    incoming_checkpoint = tune.get_checkpoint()
    if incoming_checkpoint:
        with incoming_checkpoint.as_directory() as checkpoint_dir:
            checkpoint_state = torch.load(
                Path(checkpoint_dir) / "checkpoint.pt",
                map_location=device,
                weights_only=True,
            )
        model.load_state_dict(checkpoint_state["model_state"])
        optimizer.load_state_dict(checkpoint_state["optimizer_state"])
        start_epoch = checkpoint_state["epoch"] + 1

    train_dataset, validation_dataset, _ = load_cifar10(data_dir)
    train_loader = DataLoader(
        train_dataset, batch_size=config["batch_size"], shuffle=True
    )
    validation_loader = DataLoader(
        validation_dataset, batch_size=256, shuffle=False
    )

    for epoch in range(start_epoch, max_epochs):
        model.train()
        for inputs, labels in train_loader:
            inputs = inputs.to(device, non_blocking=True)
            labels = labels.to(device, non_blocking=True)
            optimizer.zero_grad()
            loss = criterion(model(inputs), labels)
            loss.backward()
            optimizer.step()

        model.eval()
        validation_loss = 0.0
        correct_predictions = 0
        prediction_count = 0
        with torch.no_grad():
            for inputs, labels in validation_loader:
                inputs = inputs.to(device, non_blocking=True)
                labels = labels.to(device, non_blocking=True)
                outputs = model(inputs)
                validation_loss += criterion(outputs, labels).item() * labels.size(0)
                correct_predictions += (outputs.argmax(dim=1) == labels).sum().item()
                prediction_count += labels.size(0)

        metrics = {
            "loss": validation_loss / prediction_count,
            "accuracy": correct_predictions / prediction_count,
        }
        with tempfile.TemporaryDirectory() as checkpoint_dir:
            torch.save(
                {
                    "epoch": epoch,
                    "model_state": model.state_dict(),
                    "optimizer_state": optimizer.state_dict(),
                },
                Path(checkpoint_dir) / "checkpoint.pt",
            )
            tune.report(
                metrics, checkpoint=Checkpoint.from_directory(checkpoint_dir)
            )

The search samples the fully connected layer sizes, learning rate, and batch size. The ASHA scheduler can stop a trial after its first epoch when its validation loss is unlikely to compete with better trials.

search_space = {
    "first_hidden_size": tune.choice([64, 128, 256]),
    "second_hidden_size": tune.choice([32, 64, 128]),
    "learning_rate": tune.loguniform(1e-4, 1e-1),
    "batch_size": tune.choice([64, 128, 256]),
}

asha_scheduler = ASHAScheduler(
    max_t=MAX_EPOCHS,
    grace_period=1,
    reduction_factor=2,
)

trainable = tune.with_resources(
    tune.with_parameters(
        train_cifar, data_dir=DATA_DIR, max_epochs=MAX_EPOCHS
    ),
    resources={"cpu": CPUS_PER_TRIAL, "gpu": GPUS_PER_TRIAL},
)

Before you run the next cell, open the dashboard link printed by ray_init(). On the Jobs page, open the running job to inspect the trial actors, task logs, and GPU reservations. With a 0.5 GPU reservation per trial, Ray can run two trials at the same time on 1xA10 compute.

tuner = tune.Tuner(
    trainable,
    param_space=search_space,
    tune_config=tune.TuneConfig(
        metric="loss",
        mode="min",
        scheduler=asha_scheduler,
        num_samples=NUM_SAMPLES,
        max_concurrent_trials=2,
    ),
    run_config=tune.RunConfig(
        name=f"cifar10-{uuid.uuid4().hex[:8]}",
        storage_path=str(RESULTS_DIR),
        checkpoint_config=tune.CheckpointConfig(
            num_to_keep=1,
            checkpoint_score_attribute="loss",
            checkpoint_score_order="min",
        ),
    ),
)
results = tuner.fit()

if results.num_errors:
    trial_errors = "\n\n".join(str(error) for error in results.errors)
    raise RuntimeError(
        f"{results.num_errors} of {len(results)} trials failed:\n{trial_errors}"
    )

Inspect the best trial

Select the trial with the lowest reported validation loss and inspect its hyperparameters and final metrics.

best_result = results.get_best_result(metric="loss", mode="min")

print("Best hyperparameters:", best_result.config)
print(f"Validation loss: {best_result.metrics['loss']:.4f}")
print(f"Validation accuracy: {best_result.metrics['accuracy']:.2%}")

results_dataframe = results.get_dataframe()
display(
    results_dataframe[
        [
            "loss",
            "accuracy",
            "config/first_hidden_size",
            "config/second_hidden_size",
            "config/learning_rate",
            "config/batch_size",
        ]
    ].sort_values("loss")
)

Evaluate the best checkpoint

Load the selected trial's checkpoint and evaluate it on a held-out CIFAR-10 test subset.

best_model = CifarNet(
    best_result.config["first_hidden_size"],
    best_result.config["second_hidden_size"],
)
evaluation_device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
best_model.to(evaluation_device)

with best_result.checkpoint.as_directory() as checkpoint_dir:
    checkpoint_state = torch.load(
        Path(checkpoint_dir) / "checkpoint.pt",
        map_location=evaluation_device,
        weights_only=True,
    )
best_model.load_state_dict(checkpoint_state["model_state"])
best_model.eval()

_, _, test_dataset = load_cifar10(DATA_DIR)
test_loader = DataLoader(test_dataset, batch_size=256, shuffle=False)
correct_predictions = 0
prediction_count = 0
with torch.no_grad():
    for inputs, labels in test_loader:
        inputs = inputs.to(evaluation_device, non_blocking=True)
        labels = labels.to(evaluation_device, non_blocking=True)
        predictions = best_model(inputs).argmax(dim=1)
        correct_predictions += (predictions == labels).sum().item()
        prediction_count += labels.size(0)

test_accuracy = correct_predictions / prediction_count
print(f"Best checkpoint test accuracy: {test_accuracy:.2%}")

You used Ray Tune to schedule concurrent GPU trials, stop underperforming configurations with ASHA, retain checkpoints, and select a model for final evaluation. Increase NUM_SAMPLES, MAX_EPOCHS, or the dataset subset sizes when you want a more thorough search.

Example notebook

Tune a CIFAR-10 model with Ray Tune

Get notebook