Ajuster un modèle CIFAR-10 avec Ray Tune
Utilisez Ray Tune et l'algorithme ASHA (asynchronous successive halving algorithm) pour rechercher des hyperparamètres pour un classificateur d'images PyTorch sur un compute AI Runtime 1xA10 associé. Ce notebook montre comment :
- Start Ray on attached GPU compute.
- Définissez une fonction d'entraînement CIFAR-10 avec point de contrôle.
- Exécutez quatre essais Ray Tune, dont deux essais avec fractionnement de GPU en cours d’exécution simultanée.
- Sélectionnez et évaluez le meilleur essai.
Cet exemple nécessite l'environnement Databricks AI version 6 ou une version ultérieure.
Prérequis
Connectez ce notebook au compute AI Runtime GPU :
- Sélectionnez Connect en haut du notebook.
- Sélectionnez Serverless GPU .
- In the Environment side panel, set Accelerator to 1xA10 .
- Sélectionnez AI v6 comme environnement de base.
- Sélectionnez Apply , puis Confirm .
AI v6 inclut Ray, PyTorch et torchvision ; cet exemple n'installe donc pas de packages supplémentaires. La première exécution download le dataset CIFAR-10.
Initialiser Ray
Start Ray pour la session de Notebook avec ray_init(). La fonction affiche un Link vers le tableau de bord qui fonctionne par l'intermédiaire du proxy du driver Databricks.
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
Importer des bibliothèques et configurer la recherche
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"
Préparer le dataset et le modèle
Le verrou de fichier empêche des essais simultanés de download CIFAR-10 dans le même répertoire en même temps. Chaque essai utilise les mêmes sous-ensembles de formation et de validation déterministes.
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)
Définissez la fonction d'entraînement Ray Tune
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)
)
Configurer la recherche d'hyperparamètres
La recherche échantillonne la taille des couches entièrement connectées, le taux d’apprentissage et la taille de batch. L’ordonnanceur ASHA peut interrompre un essai après sa première époque lorsque sa perte de validation a peu de chances de rivaliser avec de meilleurs essais.
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
Avant d’exécuter la cellule suivante, ouvrez le Link du tableau de bord affiché par ray_init(). Sur la page Jobs , ouvrez le job en cours d’exécution pour inspecter les acteurs d’essai, les logs des tâches et les réservations de GPU. Avec une réservation de 0,5 GPU par essai, Ray peut exécuter deux essais en même temps sur un compute 1xA10.
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}"
)
Inspecter le meilleur essai
Sélectionnez l’essai présentant la perte de validation signalée la plus faible et inspectez ses hyperparamètres ainsi que ses métriques finales.
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")
)
Évaluer le meilleur point de contrôle
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%}")
Vous avez utilisé Ray Tune pour planifier des essais GPU concurrents, arrêter les configurations sous-performantes avec ASHA, conserver des points de contrôle et sélectionner un modèle pour l'évaluation finale. Augmentez NUM_SAMPLES, MAX_EPOCHS ou les tailles de sous-ensembles de datasets si vous souhaitez une recherche plus approfondie.