AIランタイムでのトレーニングパフォーマンスとレジリエンスの向上

Important

この機能は パブリック プレビュー段階です

ジョブがより多くのGPUにスケールするにつれて、ハードウェアやソフトウェアの故障の可能性は高まります。 このページでは、トレーニングランをより速く、よりフォールトトレランスにするための戦略を紹介しています:

これらのパターンにより、モデルのチェックポイントは安価なので、頻繁にチェックポイントし、安価に再開でき、GPUの有効計算能力を向上させます。

Note

serverless_gpu.data.UCVolumeDatasetserverless_gpu.data.DataLoaderserverless_gpu.data.UCVolumeWriterserverless_gpu.data.UCVolumeReaderGPU 環境5以上(サーバーレスGPU Python API 0.5.16以上)が必要です。

効率的にデータを読み込み、アイドル状態のGPU時間を最小限に抑えましょう

トレーニングステップはGPU計算と次のステップのデータ準備を重ねるべきです。 AIランタイムでは、すべてのデータアクセスがUnity Catalog経由で行われます。 Unity Catalogボリューム内のファイルベースのデータセットには serverless_gpu.data.UCVolumeDatasetを使い、FUSEマウントから各ファイルを初回アクセス時に高速ローカルキャッシュにコピーし、キャッシュされたローカルパスを得ます。

これを serverless_gpu.data.DataLoaderと組み合わせると良いでしょう。はPyTor DataLoader chのドロップインサブクラスで、サーバーレスGPUのI/Oに対応しており、GPUが計算している間にファイルを同時に取得・キャッシュします。

import serverless_gpu.data

dataset = serverless_gpu.data.UCVolumeDataset("/Volumes/my-catalog/my-schema/my-volume/data")

loader = serverless_gpu.data.DataLoader(
    dataset,
    batch_size=64,
)

for batch in loader:
    local_paths = batch  # open these immediately; see the caching note below
    ...

Warning

serverless_gpu.data.UCVolumeDatasetがもたらす道はいものです。 キャッシュは空きディスクが閾値(キャッシュファイルシステムのデフォルトの10%、 SGC_FSLAYER_MIN_FREE_DISK_BYTES 環境変数で上書き可能)を下回ると、最も最近ダウンロードしたファイルを削除します。したがって、次のアイテムを引くとパスが削除されることがあります。 同じループの繰り返しで開いたり、復号したり、コピーしたりします。 戻されたパスをリストや後で再開するような指示に保存しないでください。

パスストリームを消費する2serverless_gpu.data.UCVolumeDatasetIterableDatasetでラップしてファイルをデコードします。 ラッパーはすでにキャッシュされたローカルパスを受け取るため、解析はFUSEマウントに触れません。

from torch.utils.data import IterableDataset
from PIL import Image
import torchvision.transforms.functional as TF

class ImageDataset(IterableDataset):
    """Decodes each cached file path from UCVolumeDataset into a tensor."""

    def __init__(self, path_dataset: serverless_gpu.data.UCVolumeDataset):
        self._path_dataset = path_dataset

    def __iter__(self):
        for local_path in self._path_dataset:
            image = Image.open(local_path).convert("RGB")
            yield TF.to_tensor(image)


path_dataset = serverless_gpu.data.UCVolumeDataset("/Volumes/my-catalog/my-schema/my-volume/images")
dataset = ImageDataset(path_dataset)
loader = serverless_gpu.data.DataLoader(dataset, batch_size=64)

スケールアウト時の2つの要件があります:

  • マルチエポックトレーニングには必ず serverless_gpu.data.DataLoader を使いましょう。 persistent_workers=True時にnum_workers > 0を強制するため、各ワーカーのメモリ内キャッシュ追放トラッカーが時代を超えて存続します。 標準のPyTorch DataLoader デフォルトで毎エポックでワーカーを再フォークし、共有キャッシュディレクトリが満杯になるまでリークします。
  • すべての階級は同じ道 num_workersを通過しなければなりません。 serverless_gpu.data.UCVolumeDataset ファイルをスロット間でグローバルストライド world_size × num_workers パーティション分割します。 値が一致しないと、ファイルの重複やスキップがランク間で発生します。

torch.distributedが初期化されると、serverless_gpu.data.UCVolumeDataset反復時にランクを読み取り、ファイルをランク間で自動的に分割するため、ファイルベースのボリュームデータにはDistributedSamplerは不要です。

分散チェックポイント(DCP)付きチェックポイント

ではなくPyTorchtorch.save(DCP)を使用してください。 各ランクは自分のシャードを並列にチェックポイントディレクトリに書き込み、I/Oの全帯域幅を使い、すべての状態を1つのランクに集める際のメモリの急増を回避します。 DCPはまた、グローバルテンソルメタデータを保存するため、あるGPU数で保存されたチェックポイントを別の数で再開できます。

AIランタイムでは、 serverless_gpu.data.UCVolumeWriterserverless_gpu.data.UCVolumeReader がDCPのストレージバックエンドです。 すべてのI/Oは高速ローカルディレクトリ(/tmp、AIR GPUノード上でNVMe支援)を通じてステージ化し、Unity Catalogボリュームへのアップロードやダウンロードを行います。これは、シャードをFUSEマウントに直接書き込むよりも速いです。

import torch.distributed.checkpoint as dcp
from torch.distributed.checkpoint.state_dict import get_state_dict, set_state_dict
import serverless_gpu.data

checkpoint_path = "/Volumes/my-catalog/my-schema/my-volume/checkpoints/step_1000"

# Save
model_sd, optim_sd = get_state_dict(model, optimizer)
state_dict = {"model": model_sd, "optim": optim_sd, "step": 1000}
dcp.save(state_dict, storage_writer=serverless_gpu.data.UCVolumeWriter(checkpoint_path))

# Load
model_sd, optim_sd = get_state_dict(model, optimizer)
state_dict = {"model": model_sd, "optim": optim_sd}
dcp.load(state_dict, storage_reader=serverless_gpu.data.UCVolumeReader(checkpoint_path))
set_state_dict(
    model,
    optimizer,
    model_state_dict=state_dict["model"],
    optim_state_dict=state_dict["optim"],
)

DCPは、重みがランク間で複製される純粋なデータ並列(DDP)トレーニングでも使う価値があります。 DCPは複製された重みの重複を一つだけ書き込みつつ、各ランクの固有状態(データ位置とRNG状態、以下で説明)を保持し、後にFSDPやテンソル並列化に移行する際に必要とされるAPIと同じです。

非同期保存

同期セーブはバイトがボリュームに耐久するまでトレーニングをブロックします。 大きなチェックポイントの場合はGPUのアイドル時間です。 dcp.async_save 状態をステージングバッファに高速でコピーし、トレーニングを続ける間バックグラウンドでアップロードします。 各チェックポイントはほとんどGPU時間を消費しないため、より頻繁にチェックポイントを使えるため、中断後に失われる作業範囲が大きくなります。

非同期セーブはプロセスグループにCPUバックエンドが必要なため、 glooncclの両方で初期化してください:

import torch.distributed as dist
import torch.distributed.checkpoint as dcp
from torch.distributed.checkpoint.state_dict import get_state_dict
import serverless_gpu.data

dist.init_process_group(backend="cpu:gloo,cuda:nccl")

checkpoint_future = None

def save_async(step, model, optimizer):
    global checkpoint_future
    # Ensure the previous async save finished before starting a new one.
    if checkpoint_future is not None:
        checkpoint_future.result()

    model_sd, optim_sd = get_state_dict(model, optimizer)
    state_dict = {"model": model_sd, "optim": optim_sd, "step": step}
    writer = serverless_gpu.data.UCVolumeWriter(f"/Volumes/my-catalog/my-schema/my-volume/checkpoints/step_{step}")
    checkpoint_future = dcp.async_save(state_dict, storage_writer=writer)

最新の有効なチェックポイントから自動的に回復します

実行中にセーブ中に中断されることがあり、その場合は部分的なチェックポイントディレクトリが残ります。 serverless_gpu.data.UCVolumeWriter シャードデータファイルのアップロードが終わった後にのみ .metadata ファイルをボリュームに公開するため、 .metadata の存在はセーブ完了の信頼できるサインとなります。 リスタート時に最新の有効なチェックポイントを選択するために使ってください。

import os

def find_latest_valid(checkpoint_root):
    """Return the newest checkpoint directory that finished writing, or None."""
    candidates = sorted(
        (d for d in os.listdir(checkpoint_root) if d.startswith("step_")),
        key=lambda d: int(d.split("_")[1]),
        reverse=True,
    )
    for name in candidates:
        path = os.path.join(checkpoint_root, name)
        if os.path.exists(os.path.join(path, ".metadata")):  # save completed
            return path
    return None  # nothing valid; start fresh

レジリエントトレーニングループは、最新の有効なチェックポイントを選択し、そこから復元し、頻繁にチェックポイントを設置します。 中断後の作業損失はチェックポイント間隔の上限があるため、頻繁で安価な非同期セーブは計算量を小さく保つ。

CHECKPOINT_EVERY = 100

latest = find_latest_valid("/Volumes/my-catalog/my-schema/my-volume/checkpoints")
start_step = 0
if latest is not None:
    model_sd, optim_sd = get_state_dict(model, optimizer)
    state = {"model": model_sd, "optim": optim_sd, "step": 0}
    dcp.load(state, storage_reader=serverless_gpu.data.UCVolumeReader(latest))
    set_state_dict(
        model,
        optimizer,
        model_state_dict=state["model"],
        optim_state_dict=state["optim"],
    )
    start_step = state["step"]

for step in range(start_step, total_steps):
    train_step(...)
    if step % CHECKPOINT_EVERY == 0:
        save_async(step, model, optimizer)  # inexpensive, so run it often

このループはチェックポイントモデルと最適化状態をモデル化します。 データパイプラインの位置はまだ復元されていませんが、これは次のセクションで説明します。

データパイプラインのチェックポイント

モデルのチェックポイントはモデルと最適化器の状態を捉えますが、データセット内でのデータパイプラインの位置は記録しません。 ステップ1,900でモデルを復元したが、データローダーがデータセットの最初から再起動したとします。 再開された実行は、すでにこの時代に見られた例を再学習し、中断点付近の例をスキップすることでデータ分布に誤りなくバイアスをかけます。

正しいデータで履歴書を書くには、自分のトレーニング状態の一部としてデータセット内での位置を追跡し、履歴書にそれを復元してください。 考慮すべき点は4つあります。

サンプルやシャードのオフセットを追跡する

チェックポイントステートの命令で、グローバルサンプルインデックス、バッチカウント、消費されたシャードIDのリストを使ってエポックの進行度を記録し、再開時にその位置にスキップします。 これにより、データローダーに内部状態のシリアライズを頼るのではなく、あなたの明確なコントロール下にデータ位置を保持できます。

# Include the data position in the checkpoint state dict:
state_dict = {
    "model": model_sd,
    "optim": optim_sd,
    "step": step,
    "epoch": epoch,
    "samples_seen": samples_seen,  # your own counter, advanced each batch
}

決定論的サンプラーを持つマップスタイルデータセットでは、再開時にこのエポックで既に消費されたバッチをスキップしてください。 サンプラーの順序は特定の (seed, epoch) に対して決定的であるため(パイプライン を決定的にする項目を参照)、早送りは正確な位置を再現します:

resume_batch = state["samples_seen"] // batch_size

for epoch in range(start_epoch, num_epochs):
    sampler.set_epoch(epoch)
    for batch_idx, batch in enumerate(loader):
        # On the resumed epoch only, skip batches already processed.
        if epoch == start_epoch and batch_idx < resume_batch:
            continue
        train_step(batch)
        samples_seen += batch_size

シャード化されたストリーミングデータセットの場合は、完了したシャードのセットを追跡し、残ったシャードだけを再開実行に渡します。 これにより、中断点に到達するために何度もバッチを再生する必要がありません:

# Filter the shard list down to work not yet done, then build the loader from it.
remaining = [s for s in all_shards if s not in state["completed_shards"]]
dataset = ShardDataset(remaining)

カスタムデータセットの内部状態をチェックポイントする

自分でデータセットを書く場合は、そのデータセットにシリアライズや位置復元のメソッドを与え、その状態をチェックポイントに折りたたみ入れましょう。 これにより、履歴ロジックが反復ロジックの隣に位置し、データセットはファストフォワードに必要なもの(現在のシャード、その中のオフセット、バッファ内容のシャッフルなど)を正確に把握し、トレーニングループが外部カウンターから再構築するのを防ぎます。

from torch.utils.data import IterableDataset

class ResumableShardDataset(IterableDataset):
    """A streaming dataset that can checkpoint and restore its own position."""

    def __init__(self, shards):
        self._shards = shards
        self._shard_idx = 0      # position advanced during __iter__
        self._offset = 0

    def state_dict(self):
        return {"shard_idx": self._shard_idx, "offset": self._offset}

    def load_state_dict(self, state):
        self._shard_idx = state["shard_idx"]
        self._offset = state["offset"]

    def __iter__(self):
        for i in range(self._shard_idx, len(self._shards)):
            self._shard_idx = i
            for j, example in enumerate(self._read_shard(self._shards[i])):
                if j < self._offset:
                    continue  # skip examples already consumed from this shard
                self._offset = j + 1
                yield example
            self._offset = 0
# Save and restore the dataset position with the rest of the checkpoint state.
state_dict["dataset"] = dataset.state_dict()
# On resume:
dataset.load_state_dict(state["dataset"])

時代境界からの再スタート

スキップアヘッドが非現実的の場合は、エポック境界でのみチェックポイントし、次のエポック開始時に再開します。 その後、中断期間中の各エポックごとにすべてのインスタンスが正確に一度ずつ見られ、失敗ごとに最大1エポクの進行状況を失います。 これは、エポックが故障率に対して短い場合に最も単純です。

パイプラインを決定論的にしましょう

どちらの戦略も、保存状態からパイプラインが再現可能である場合にのみ正しいデータで再開されます。 シャッフルやオーグメンテーションはRNGから引き起こすので、シードしてチェックポイントでその状態を持ち続けましょう。 そうでなければ、再起動後のシャッフルや増幅の順序が中断前の順序と一致せず、スキップアヘッドは誤ったサンプルにオフセットを付けてしまいます。

パイプラインの途中で再開するには、シャッフルや拡張を駆動するRNG自体がチェックポイント可能でなければなりません。 シーディングだけでは最初から順序が再現されますが、中断された地点は再現できません。 内部状態をシリアライズできるRNGオブジェクトを使い、その状態をチェックポイントして各RNGが前回のまま続きます。 エポック開始時にのみ再シドされるグローバルRNGに頼ると、エポック開始時と同じ抽選を再生し、中間エポックのスキップアヘッドオフセットと一致しません。

すべてのランダム性の源をまきましょう:

import random
import numpy as np
import torch

def seed_everything(seed: int):
    random.seed(seed)
    np.random.seed(seed)
    torch.manual_seed(seed)
    torch.cuda.manual_seed_all(seed)

モデルと一緒にRNG状態を保存・復元することで、拡張やシャッフルのシーケンスがシームレスに続きます:

# Save
state_dict["rng"] = {
    "python": random.getstate(),
    "numpy": np.random.get_state(),
    "torch": torch.get_rng_state(),
    "cuda": torch.cuda.get_rng_state_all(),
}

# Load
rng = state["rng"]
random.setstate(rng["python"])
np.random.set_state(rng["numpy"])
torch.set_rng_state(rng["torch"])
torch.cuda.set_rng_state_all(rng["cuda"])

ステートフルデータローダーの代わりに DistributedSampler を使う場合は、各エポックの開始時に sampler.set_epoch(epoch) を呼び出します。 シャッフルは (seed, epoch) の決定論的関数であるため、エポックカウンターを復元すると正確な置換が再現されます。

for epoch in range(start_epoch, num_epochs):
    sampler.set_epoch(epoch)  # deterministic reshuffle per epoch
    for batch in loader:
        ...

Note

データパイプラインの正確性のためには、データの順序と増強ストリームが再現可能であることが必要であり、上記のシーディングとRNGチェックポイントがそれを提供しています。 一般的にビットごとに同一のフォワードパスは必要ありません。 torch.use_deterministic_algorithms(True) 決定論的カーネルを強制しますが、スループットを低下させる可能性があり、すべての操作をカバーするわけではありません。