Skip to main content

Tune a CIFAR-10 model with Ray Tune

Open in Databricks

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.

Python
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.

Python
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.

Python
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.

Python
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.

Python
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.

Python
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.

Python
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.

Python
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