TT Lab
はじめる
学ぶ 学習パス コース

MiniMind — 小さな言語モデルを最初から最後まで自分で学習する

再開ファイルには重みのほかに3つが入る

TT Labで続きを見る

一言でいうと

学習は途中で途切れます。セッションが終わり、ノードが再起動し、コストのために止まります。止まった位置から一度に最後まで進んだのとまったく同じにつなぐには、重みだけでは足りません。MiniMindのlm_checkpointは、重みとは別に再開ファイルを置き、そこにオプティマイザーの状態・学習ステップ数・(分散学習なら)プロセス数を一緒に入れます。このモジュールでは、60学習ステップを一度に学習した結果と、30学習ステップで止めてからつないだ結果が、1ビットも違わないようにして、1つずつ抜いたときにどれだけ離れるかを測ります。

なぜ必要なのか

重みだけを保存してから続きを学習すると、結果が静かに変わります。エラーは出ません。損失曲線も、それらしくつながります。ところが、2つの点がずれています。

学習率のスケジュールは、別に保存しなくてもかまいません。学習ステップ番号で計算する関数(get_lr)なので、学習ステップ数さえあればよいのです。分散学習でGPUの数が変わると、1学習ステップが見るデータの量が変わるので、MiniMindは、保存したworld_sizeと現在の数で、学習ステップ数を換算します。

どう動くのか

MiniMindの保存は、2つのファイルを書きます。

state_dict = {k: v.half().cpu() for k, v in raw_model.state_dict().items()}
torch.save(state_dict, ckp_path + ".tmp");  os.replace(ckp_path + ".tmp", ckp_path)
resume = {"model": state_dict, "optimizer": optimizer.state_dict(),
          "epoch": epoch, "step": step, "world_size": ..., "wandb_id": ...}
torch.save(resume, resume_path + ".tmp");  os.replace(resume_path + ".tmp", resume_path)

一時ファイル → 名前の変更では、保存している途中でプロセスが死ぬと、半分だけ書かれたファイルが残ります。そのファイルが元の名前を上書きしていたら、無事だった以前のチェックポイントまで失います。.tmpに最後まで書いてからos.replaceで一度に入れ替えれば、同じファイルシステムの中ではアトミックに入れ替わるので、いつ死んでも、完全なファイルが1つ残ります。

半精度では、重みを.half()で保存して、ファイルを半分に減らします。読み込んでfp32に広げると、ロジットが少し変わります(このコースのモデルで最大0.003前後)。推論には問題ありませんが、再開ファイルの重みまで半精度なら、続きを行った学習が、一度に行ったものとビット単位で同じにはなれません。このラボは、比較のために、再開ファイルにfp32の重みを入れます。

安全な読み込みでは、torch.loadは、Pythonのオブジェクトを復元するpickleを使います。他人からもらったチェックポイントをそのまま読むと、中のコードが実行されることがあります。weights_only=Trueで読むと、テンソル・数・文字列・dictだけを復元します。このラボの再開ファイル(乱数生成器の状態を含む)も、そのまま読めます。

現場での姿

プリエンプティブル(spot)インスタンスで学習すると、数時間ごとにノードを奪われます。再開が正確でないと、毎回少しずつ違う学習になり、最終モデルを再現できず、問題が起きても原因を絞り込めません。そのため、「続きを行ったもの = 一度に行ったもの」を、短い学習で先に確認しておきます。保存がアトミックでないと、逆方向の事故が起きます。保存の途中でノードが死んで、最新のチェックポイントが壊れ、その前のものはすでに上書きされてなくなっているのです。

MiniMindのオリジナルとこのコースの違い

MiniMindは、学習ステップごとに新しく選ぶ代わりに、エポック単位の順序を使います。エポックごとにsetup_seed(seed + epoch)でtorch.randpermの順序を作り、再開したら同じ順序を作り直して、SkipBatchSamplerですでに見たバッチの数だけ飛ばします。このコースのループは、学習ステップごとに任意の位置を選ぶので、同じことを、バッチの乱数生成器の状態(get_state・set_state)を保存・復元することで行います。方法は違っても、守るものは同じです。「次に見るデータ」が、止まった位置からつながることです。また、MiniMindは重みを半精度で保存するので、続きを行った学習がビット単位で同じにはなりません。このコースは、その違いまでなくして確認するために、再開ファイルにfp32の重みを入れます。

次のラボですること

60学習ステップを一度に学習した結果を基準にして、30学習ステップで止めて再開ファイルをアトミックに保存し、続けて60学習ステップまで進んで、重みが同じかを見ます。オプティマイザーの状態を捨ててつないだもの、バッチの順序を巻き戻したものが、どれだけ離れるかを測り、半精度で保存したときのサイズとロジットの差、再開ファイルでオプティマイザーの状態が占める割合を数えます。