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.
This example requires the Databricks AI environment version 6 or above.
Prerequisites
Connect this notebook to AI Runtime GPU compute:
- Select Connect at the top of the notebook.
- Select Serverless GPU.
- In the Environment side panel, set Accelerator to 1xA10.
- Select AI v6 as the base environment.
- 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
Import libraries and configure the search
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)
)
Configure the hyperparameter search
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},
)
Run and monitor the search
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.