Ajuste um modelo CIFAR-10 com o 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:
- Inicie o Ray no compute de GPU conectado.
- Defina uma função de treinamento do CIFAR-10 com ponto de verificação.
- Execute quatro testes do Ray Tune, com dois testes de GPU fracionária sendo executados de forma concorrente.
- Selecione e avalie o melhor teste.
Este exemplo requer o ambiente do Databricks AI versão 6 ou acima.
Pré-requisitos
Conecte este notebook ao AI Runtime GPU compute:
- Select Connect at the top of the notebook.
- Selecione Serverless GPU .
- No painel lateral Ambiente , defina Acelerador como 1xA10 .
- Select AI v6 as the base environment.
- Selecione Apply e, em seguida, selecione Confirm .
O AI v6 inclui Ray, PyTorch e torchvision, portanto este exemplo não instala pacotes adicionais. A primeira execução faz o download do dataset CIFAR-10.
Inicializar o Ray
Começar o Ray para a sessão do notebook com ray_init(). A função exibe um link de painel que funciona através do proxy de driver do 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
Importar bibliotecas e configurar a pesquisa
O exemplo limita cada teste a um subconjunto do CIFAR-10 e a cinco épocas de treinamento para reduzir o Runtime de pesquisa. Cada teste reserva metade da GPU A10, o que permite que o Ray programe até duas tentativas simultaneamente.
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 o dataset e o modelo
O bloqueio de arquivo impede que tentativas concorrentes de download do CIFAR-10 para o mesmo diretório ao mesmo tempo. Cada tentativa usa os mesmos subconjuntos determinísticos de treinamento e validação.
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)
Defina a função de treinamento do Ray Tune
O Ray Tune chama essa função para cada configuração de hiperparâmetro amostrada. No final de cada época, a função relata métricas de validação e um ponto de verificação. O ASHA usa a perda relatada para interromper antecipadamente as tentativas com baixo desempenho.
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 programador 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},
)
Executar e monitorar a busca
Antes de executar a próxima célula, abra o link do painel impresso por ray_init(). Na página Jobs , abra o job em execução para inspecionar os atores de teste, os logs de tarefas e as reservas de GPU. Com uma reserva de 0,5 GPU por teste, o Ray pode executar dois testes ao mesmo tempo em um 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}"
)
Inspecionar o melhor teste
Selecione o teste com a menor perda de validação relatada e inspecione seus hiperparâmetros e métricas finais.
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")
)
Avalie o melhor ponto de verificação
Carregue o checkpoint do teste selecionado e avalie-o em um subconjunto de teste CIFAR-10 retido.
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.