メインコンテンツまでスキップ

Ray Tune を使用した CIFAR-10 モデルのチューニング

Open in Databricks

Ray Tune と非同期逐次半減アルゴリズム (ASHA) を使用して、アタッチされた 1xA10 AI ランタイムコンピュート上の PyTorch 画像分類器のハイパーパラメータを検索します。このノートブックでは、次の方法について説明します。

  • アタッチされた GPU コンピュートで Ray を起動します。
  • チェックポイントが設定されたCIFAR-10トレーニング関数を定義します。
  • 2つの部分GPUトライアルを同時に実行しながら、4つのRay Tuneトライアルを実行します。
  • 最適なトライアルを選択して評価します。
注記

この例では、Databricks AI 環境バージョン 6 以上が必要です。

前提条件​

このノートブックをAIランタイムGPUコンピュートに接続します:

  1. ノートブックの上部にある [接続] を選択します。
  2. Serverless GPU を選択します。
  3. Environment サイドパネルで、 Accelerator を 1xA10 に設定します。
  4. AI v6 を基本環境として選択します。
  5. [適用] を選択し、続いて [確認] を選択します。

AI v6にはRay、PyTorch、およびtorchvisionが含まれているため、この例では追加のパッケージはインストールされません。最初のランでCIFAR-10データセットをdownloadします。

Ray の初期化​

ray_init() を使用してノートブックセッションで Ray を起動します。この関数は、Databricks ドライバー プロキシを介して機能するダッシュボード Link を出力します。

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

ライブラリのインポートと検索の構成​

検索ランタイムを短縮するため、この例では各試行を CIFAR-10 のサブセットと 5 回のトレーニングエポックに制限しています。各試行は A10 GPU の半分を確保するため、Ray は最大 2 つの試行を同時にスケジュールできます。

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"

データセットとモデルを準備する​

ファイル ロックにより、複数の並列トライアルが同時に CIFAR-10 を同じディレクトリに download するのを防ぎます。すべてのトライアルで同じ決定的トレーニングおよび検証サブセットが使用されます。

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)

Ray Tune トレーニング関数を定義する​

Ray Tune は、サンプリングされたハイパーパラメータ設定ごとにこの関数を呼び出します。各エポックの終了時に、関数は検証メトリクスとチェックポイントを報告します。ASHA は、報告された損失を使用して、パフォーマンスの低い試行を早期に停止します。

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)
)

ハイパーパラメータ探索を構成する​

探索では、全結合層のサイズ、学習率、バッチサイズをサンプリングします。ASHA スケジューラは、検証ロスが優れたトライアルと競合する可能性が低い場合、最初のエポックの後にトライアルを停止できます。

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},
)

検索の実行と監視​

次のセルを実行する前に、ray_init() によって出力されたダッシュボードのLinkを開いてください。 [ジョブ] ページで実行中のジョブを開き、トライアル アクター、タスク Logs、および GPU 予約を確認します。トライアルあたり 0.5 個の GPU を予約すると、Ray は 1 台の A10 コンピュートで同時に 2 つのトライアルを実行できます。

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}"
)

最適なトライアルの調査​

報告された検証損失が最も低い試行を選択し、そのハイパーパラメータと最終メトリクスを確認します。

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")
)

最適なチェックポイントを評価する​

選択したトライアルのチェックポイントをロードし、ホールドアウトされた CIFAR-10 のテスト サブセットで評価します。

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%}")

Ray Tune を使用して並列 GPU 試行をスケジュールし、ASHA でパフォーマンスの低い構成を停止し、チェックポイントを保持して、最終評価用のモデルを選択しました。より徹底した検索を行うには、NUM_SAMPLES、MAX_EPOCHS、またはデータセットのサブセットサイズを増やします。

ノートブックの例​

Ray Tune を使用した CIFAR-10 モデルのチューニング