0
0

Delete article

Deleted articles cannot be recovered.

Draft of this article would be also deleted.

Are you sure you want to delete this article?

JAX/Flax の学習ループを実装する上で学んだことまとめ

0
Posted at

はじめに

前回までに、JAX の基本と、Flax でのモデル作成・PyTorch との重み変換について書きました。

Flax でモデルを作るところまでは前回の記事で書きました、それを学習させて重みを保存するところまで行ったので、その過程で学習したことについて記します。この記事では、MNIST を題材に、学習ループを書き、一部のパラメータを凍結し、checkpoint を保存・復元するところまでを通しで書きます。途中で実際に踏んだエラーも載せます。

動作確認した環境は、jax 0.11.1、flax 0.12.9、optax 0.2.8 です。

学習ループの全体像

PyTorch では loss.backward()optimizer.step() を呼ぶと、モデルと optimizer が内部状態を持ったまま更新されます。JAX は関数型なので、状態を明示的に受け取り、新しい状態を返す関数として書きます。

更新のたびに、paramsopt_state が新しいものに置き換わります。この2つと step を1つにまとめて持ち回るのが、次に出てくる TrainState です。

TrainState

flax.training.train_state.TrainState は、paramsopt_statestep と、apply_fn(モデルの apply)・tx(optimizer)をまとめたデータクラスです。

import flax.linen as nn
import jax
import jax.numpy as jnp
import numpy as np
import optax
from flax.training import train_state


class MLP(nn.Module):
    @nn.compact
    def __call__(self, x):
        x = x.reshape((x.shape[0], -1))
        x = nn.relu(nn.Dense(128, name="hidden")(x))
        return nn.Dense(10, name="head")(x)


def create_state(model, key, tx):
    params = model.init(key, jnp.zeros((1, 28, 28, 1)))["params"]
    return train_state.TrainState.create(apply_fn=model.apply, params=params, tx=tx)

state.apply_gradients(grads=grads) を呼ぶと、次のことをまとめてやってくれます。

新しい TrainState が返るだけで、元の state は書き換わらない点に注意です。

optax で optimizer を組む

optax は、勾配の変換を部品として組み合わせて使います。今回は「勾配のクリップ → Adam」を chain でつなぎ、学習率には cosine decay のスケジュールを渡します。

batch_size = 128
steps_per_epoch = 60000 // batch_size
epochs = 3

schedule = optax.cosine_decay_schedule(1e-3, steps_per_epoch * epochs, alpha=0.1)
tx = optax.chain(optax.clip_by_global_norm(1.0), optax.adam(schedule))

model = MLP()
state = create_state(model, jax.random.key(0), tx)

alpha=0.1 は、最終的な学習率が初期値の10%になる、という意味です。optax.adam に固定値ではなく関数(schedule)を渡すと、内部で step を数えて学習率を変えてくれます。この設定は、後の checkpoint の復元で問題になります。

1 step の更新関数

@jax.jit
def train_step(state, images, labels):
    def loss_fn(params):
        logits = state.apply_fn({"params": params}, images)
        loss = optax.softmax_cross_entropy_with_integer_labels(logits, labels).mean()
        return loss, logits

    (loss, logits), grads = jax.value_and_grad(loss_fn, has_aux=True)(state.params)
    accuracy = (logits.argmax(-1) == labels).mean()
    return state.apply_gradients(grads=grads), loss, accuracy

jax.value_and_grad で、損失と勾配を同時に求めます。損失以外の値(ここでは logits)も返したいときは has_aux=True にします。また、@jax.jit でこの関数全体がコンパイルされます。

jit に関数を渡すと怒られる

評価用の関数を書くとき、apply_fn をそのまま引数に渡してしまうと、次のエラーになります。

@jax.jit
def eval_step(params, apply_fn, images, labels):   # NG
    ...
TypeError: Cannot interpret value of type <class 'method'> as an abstract array;
it does not have a dtype attribute

jit の引数は「配列の PyTree」でなければならず、関数は配列ではないので受け取れません。対処は2つあります。

  • static_argnums で「コンパイル時に固定する値」として扱う(関数が変わるたびに再コンパイルされます)
  • TrainState ごと渡す。apply_fnTrainState の中で「PyTree の葉ではない値」として扱われるので、そのまま通ります
@jax.jit
def eval_step(state, images, labels):
    logits = state.apply_fn({"params": state.params}, images)
    return (logits.argmax(-1) == labels).sum()

shape が変わると再コンパイルされる

jit は、引数の shape と dtype ごとにコンパイルします。最後のバッチだけ端数が出ると、そこで再コンパイルが走って遅くなります。学習ループでは、端数のバッチは捨てる(drop_last)のが簡単です。

学習ループを回す

データは MNIST を使います。

from torchvision import datasets


def load_mnist():
    def to_arrays(dataset):
        x = dataset.data.numpy().astype(np.float32)[..., None] / 255.0
        y = dataset.targets.numpy().astype(np.int32)
        return x, y

    train = datasets.MNIST("data", train=True, download=True)
    test = datasets.MNIST("data", train=False, download=True)
    return to_arrays(train), to_arrays(test)


(train_x, train_y), (val_x, val_y) = load_mnist()

エポックごとにシャッフルしてから、バッチを1つずつ train_step に渡します。

key = jax.random.key(1)
for epoch in range(epochs):
    key, shuffle_key = jax.random.split(key)
    order = np.asarray(jax.random.permutation(shuffle_key, len(train_x)))
    for i in range(steps_per_epoch):
        idx = order[i * batch_size : (i + 1) * batch_size]
        state, loss, acc = train_step(state, train_x[idx], train_y[idx])

    correct = eval_step(state, val_x, val_y)
    print(f"epoch={epoch + 1} step={int(state.step)} "
          f"loss={float(loss):.3f} val_acc={float(correct) / len(val_x):.3f}")

ベースラインと比べる

検証の精度は、単体で見ても良いのか悪いのか分かりません。最頻クラスを常に答える場合の精度(ベースライン)と比べる習慣をつけておくと、「ほぼ学習できていない」状態に早く気づけます。

baseline = np.bincount(val_y).max() / len(val_y)
print(f"最頻クラスのベースライン: {baseline:.3f}")

一部のパラメータだけ学習する

事前学習した部分を固定して、出力層だけ学習したい場合があります。optax の multi_transform を使うと、パラメータのグループごとに別の optimizer を割り当てられます。

params = model.init(jax.random.key(2), jnp.zeros((1, 28, 28, 1)))["params"]

labels = jax.tree_util.tree_map_with_path(
    lambda path, _: "train" if path[0].key == "head" else "frozen", params
)
tx_freeze = optax.multi_transform(
    {"train": optax.adam(1e-2), "frozen": optax.set_to_zero()}, labels
)
state = train_state.TrainState.create(apply_fn=model.apply, params=params, tx=tx_freeze)

set_to_zero は、更新量を常に0にする変換です。勾配自体は計算されますが、パラメータには反映されません。

凍結できているかをテストで確かめる

「凍結したつもり」で学習が進んでいなかった、というのは避けたいところです。1 step 更新して、更新前後を比べるテストを1つ書いておくと安心できます。

new_state, _, _ = train_step(state, train_x[:128], train_y[:128])

hidden_same = jnp.array_equal(state.params["hidden"]["kernel"],
                              new_state.params["hidden"]["kernel"])
head_changed = not jnp.array_equal(state.params["head"]["kernel"],
                                   new_state.params["head"]["kernel"])
assert hidden_same and head_changed

checkpoint の保存と復元

flax.serialization で保存する

serialization.to_bytes でバイト列にして、ファイルに書き出せます。

from pathlib import Path
from flax import serialization

Path("ckpt.msgpack").write_bytes(serialization.to_bytes(state))

復元は、同じ構造の TrainState を用意して、そこへ値を流し込みます。

template = create_state(model, jax.random.key(3), tx)   # 保存時と同じ tx
state = serialization.from_bytes(template, Path("ckpt.msgpack").read_bytes())

学習率スケジュールの有無で復元に失敗する

「optax で optimizer を組む」の節で、optax.adam にスケジュールを渡しました。このとき opt_state には、step を数えるための count が含まれます。この opt_state を、固定の学習率で作った TrainState へ復元しようとすると失敗します。

tx_const = optax.chain(optax.clip_by_global_norm(1.0), optax.adam(1e-3))
template = create_state(model, jax.random.key(3), tx_const)
serialization.from_bytes(template, data)
ValueError: The field names of the state dict and the named tuple do not match,
got {'count'} and set() at path ./opt_state/1/1

推論のためだけに読み込みたいときや、別の optimizer で再開したいときに起こりやすいエラーです。対処は2つあります。

対処1: ダミーのスケジュールで、保存時と同じ構造の template を作る

スケジュールの中身は復元で上書きされるので、構造だけ合っていれば問題ありません。

dummy = optax.chain(
    optax.clip_by_global_norm(1.0),
    optax.adam(optax.cosine_decay_schedule(1.0, 1)),
)
template = create_state(model, jax.random.key(3), dummy)
state = serialization.from_bytes(template, data)

対処2: params だけ取り出す

optimizer の状態が不要なら、バイト列を辞書として読んで、params だけ使います。template を作らないので、optimizer の構造に依存しません。

raw = serialization.msgpack_restore(data)
params = raw["params"]

orbax を使う場合

flax.serialization は手軽ですが、ファイルの管理(世代・非同期保存など)は自分で書く必要があります。そのあたりまで含めて任せたい場合は、標準的なライブラリの orbax-checkpoint を使います。

import orbax.checkpoint as ocp

checkpointer = ocp.StandardCheckpointer()
checkpointer.save("/abs/path/params", state.params)
checkpointer.wait_until_finished()

restored = checkpointer.restore("/abs/path/params", state.params)

保存先は、絶対パスで指定するのが無難です(相対パスでの動作は試していません)。

今回の目的は、学習した重みを保存して、前回の重み変換で PyTorch へ渡すことでした。その場合は、params だけを flax.serialization で保存する形で十分でした。

GPU メモリ不足と NaN のデバッグ

GPU メモリ不足(OOM)

JAX は、既定で GPU メモリの大半を最初に確保します。他のプロセスと同じ GPU を使うときは、環境変数で調整します。

XLA_PYTHON_CLIENT_PREALLOCATE=false      # 必要になったぶんだけ確保する
XLA_PYTHON_CLIENT_MEM_FRACTION=0.5       # 使ってよい割合を制限する

もう1つ、評価で全データを一度に順伝播すると OOM になることがあります。学習は小さなバッチで回しているのに、評価だけ全件を渡してしまうケースです。jax.lax.map で分割して評価します。

@jax.jit
def predict_all(params, x):
    chunk = 256
    chunks = x.reshape((len(x) // chunk, chunk) + x.shape[1:])
    logits = jax.lax.map(lambda c: model.apply({"params": params}, c), chunks)
    return logits.reshape((len(x), -1))

この例は、len(x)chunk で割り切れる場合のものです。割り切れない場合は、端数を別に処理するか、パディングしてください。

分割しても予測は変わらないはずですが、GPU の行列積はバッチの形が変わると、ごくわずかに丸めが変わります。一括計算と比べるときは、allclose の許容誤差を 1e-5 ではなく 1e-3 くらいにしておきます。手元の確認では、最大で 1.7e-4 ずれました(argmax は全件一致でした)。

NaN の発生源を特定する

jax.config.update("jax_debug_nans", True)

有効にすると、NaN を生成した演算の時点でエラーになります。jit の中の値を確認したいときは、print ではなく jax.debug.print を使います。

jax.debug.print("loss={loss}", loss=loss)

おわりに

JAX の学習ループは、paramsopt_state を持ち回る関数として書きます。TrainState がそれをまとめてくれて、optax は chain で部品をつなぎ、一部だけ学習したいときは multi_transform を使います。

jit の引数は配列の PyTree なので、関数を渡したいときは TrainState ごと渡します。checkpoint の復元には保存時と同じ構造の template が必要で、学習率スケジュールの有無でも構造が変わるので注意が必要です。評価はベースラインと比べ、全件を一度に流さず分割して行います。

保存した params は、前回の記事の重み変換で PyTorch へ渡せます。

0
0
0

Register as a new user and use Qiita more conveniently

  1. You get articles that match your needs
  2. You can efficiently read back useful information
  3. You can use dark theme
What you can do with signing up
0
0

Delete article

Deleted articles cannot be recovered.

Draft of this article would be also deleted.

Are you sure you want to delete this article?