MiniMind — 小さな言語モデルを最初から最後まで自分で学習する
止めて再開しても一度に行ったものと同じに — 再開ファイルに何を入れるか
目標
60学習ステップの学習を一度に行った結果と、30学習ステップで再開ファイル(重み・オプティマイザーの状態・学習ステップ数・バッチの乱数状態)をアトミックに保存してから続けた結果が、1ビットも違わないようにします。オプティマイザーの状態やデータの位置を抜くとどれだけ離れるか、半精度の保存が何を変えるかを測ります。
なぜ重要なのか
重みだけを保存して続きを学習すると、エラーなしで結果が変わります。AdamWは、パラメーターごとに2つのモーメントを記憶しているので、その記憶を失うと、つないだ直後の数学習ステップが、まったく違う大きさで動きます。エポックの途中でバッチの順序を巻き戻すと、前の部分を2回見て、後ろの部分は見られません。このような違いは、損失曲線にほとんど表れないので、気づかずに通り過ぎやすいです。
MiniMindのlm_checkpointは、そのため、重みと再開ファイルを別に置き、どちらも一時ファイルに書いてから名前を変えます。このラボは、その設計が正確に何を守っているのかを、一度に行った結果と比べて確認します。
ステップ
- スクリプト(/root/mm/ckpt/run.py)を書き、60学習ステップ(シード7・バッチ8・長さ128・lr 3e-3)を一度に学習して、結果を残してください(重み: /root/mm/ckpt/straight.pth、記録(
step,train_loss): /root/mm/ckpt/straight_log.csv)。 - 同じ学習を30学習ステップで止めて、再開ファイルを
.tmpに書いてからos.replaceで入れ替えて保存してください(model・optimizer・step・gen_state。出力先: /root/mm/ckpt/resume.pth)。 - 再開ファイルから31学習ステップ目以降を60学習ステップ目まで続けて、結果を残してください(重み: /root/mm/ckpt/resumed.pth、記録: /root/mm/ckpt/resumed_log.csv)。ステップ1と同じになる必要があります。
- オプティマイザーの状態だけを捨ててつないだ結果を残してください(出力先: /root/mm/ckpt/noopt.pth)。
- バッチの乱数状態だけを捨てて(最初のシードでやり直して)つないだ結果を残してください(出力先: /root/mm/ckpt/noskip.pth)。
straight.pthを半精度で保存し、2つのファイルのサイズとロジットの差を書いてください(半精度の保存先: /root/mm/ckpt/straight_fp16.pth、サイズとロジットの差の出力先: /root/mm/ckpt/fp16.json)。- 再開ファイルで、重みとオプティマイザーの状態の要素数を数えて書いてください(出力先: /root/mm/ckpt/optim.json)。
## 무엇을 저장하나## 빠뜨리면## 반정밀도の3つのセクションを書き、ステップ7の倍率とステップ6のfp16ファイルのサイズ(バイト)を入れてください(出力先: /root/mm/ckpt/report.md)。韓国語の見出しは、順に「何を保存するのか」「抜いたら」「半精度」という意味です。
参考
- 60学習ステップは、ノードで5秒前後です。1つのスクリプトに
--stop-at・--resume・--no-optim・--no-skipを置いておけば、ステップ3–5が1行ずつで済みます。 - バッチを
torch.randint(..., generator=g)で選ぶなら、g.get_state()・g.set_state()で「次に選ぶバッチ」を保存・復元できます。MiniMindのSkipBatchSamplerがしていることを、この1行が行います。 - 再開ファイルは
torch.load(경로, weights_only=True)で読んでください(プレースホルダーはパスです)。他人からもらったチェックポイントをそのまま読むと、中のPythonコードが実行されることがあります。 - よくある間違い: モデルを作る前にシードを固定せず、再開した側の初期化が変わること(すぐに上書きされるので問題なさそうに見えますが、生成器の状態がずれます)、学習率のスケジュールを、つないだ学習ステップを基準に0から数え直すこと。
- 原典: trainer_utils.py — lm_checkpoint・SkipBatchSampler · PyTorch — Saving and Loading · torch.load weights_only · os.replace
一度に60学習ステップ学習する(基準)
スクリプト(/root/mm/ckpt/run.py)を書いて、シード7・バッチ8・長さ128・lr 3e-3(MiniMindのコサイン、全60学習ステップが基準)で、60学習ステップを一度に学習し、結果を残してください(重み(state_dict): /root/mm/ckpt/straight.pth、記録(step,train_loss、1–60): /root/mm/ckpt/straight_log.csv)。
前のモジュールの学習ループと同じです。今回は、損失を小数点以下6桁まで書いておいてください。ステップ3でつないだ記録と、1桁ずつ比べます。
30学習ステップでアトミックに保存する
同じ学習を30学習ステップで止め、{"model": state_dict, "optimizer": optimizer.state_dict(), "step": 30, "gen_state": 배치 생성기.get_state()}(プレースホルダーはバッチの生成器です)を、一時ファイル(/root/mm/ckpt/resume.pth.tmp)に書いてから、os.replaceで入れ替えて保存してください(出力先: /root/mm/ckpt/resume.pth)。学習率は、やはり全60学習ステップを基準に計算します。
一時ファイルに最後まで書いてから名前を変えれば、保存の途中で死んでも、完全なファイルが1つ残ります。MiniMindのlm_checkpointが、ckp_tmp → os.replaceで行っていることです。.tmpが残っていたら、名前の変更を忘れています。
続けて60学習ステップまで進める(基準と一致させる)
新しいプロセスでresume.pthを読み、重み・オプティマイザーの状態・バッチ生成器の状態を復元して、31学習ステップから60学習ステップまで学習し、結果を残してください(重み: /root/mm/ckpt/resumed.pth、記録(31–60): /root/mm/ckpt/resumed_log.csv)。
モデルとオプティマイザーを、ステップ1とまったく同じように作ってから、load_state_dictで上書きし、g.set_state(…)で生成器を元に戻します。採点ツールは、2つのチェックポイントのすべてのテンソルを比べて、最大の差が1e-6以下かを見ます。同じノードで回せば、0になります。
オプティマイザーの状態を捨てたらどうなるかを見る
再開ファイルから重みとバッチ生成器の状態だけを復元し、オプティマイザーは新しく作ったまま、31–60学習ステップをつないで、結果を残してください(出力先: /root/mm/ckpt/noopt.pth)。
AdamWの1学習ステップの大きさは、勾配を、その移動平均の平方根で割った値に比例します。記憶が0から再び始まると、つないだ直後の数学習ステップが、まったく違う大きさで動きます。採点ツールは、この結果が基準と確実に違うかを見ます。
バッチの順序を巻き戻したらどうなるかを見る
再開ファイルから重みとオプティマイザーの状態だけを復元し、バッチ生成器は最初のシード(7)で作り直したまま、31–60学習ステップをつないで、結果を残してください(出力先: /root/mm/ckpt/noskip.pth)。
生成器を最初の状態にしておくと、31学習ステップ目に、1学習ステップ目のバッチをもう一度見ることになります。エポックの途中で再開するとき、MiniMindがSkipBatchSamplerで、すでに見たバッチを飛ばす理由です。
半精度で保存したらどうなるかを見る
straight.pthのすべてのテンソルを.half()に変えて保存し(保存先: /root/mm/ckpt/straight_fp16.pth)、2つのファイルのサイズ(バイト)と、/opt/mm/ref/val.npyの先頭128トークンに対するロジットの最大絶対差を、fp32_bytes・fp16_bytes・max_logit_diffとして書いてください(出力先: /root/mm/ckpt/fp16.json)。
mmkit.load_modelは、fp16の重みもfp32に広げて入れます。ファイルは半分近くに減りますが(テンソル以外に名前のようなものも入っているので、ちょうど半分ではありません)、ロジットは小数点以下3桁ほどのところで変わります。
再開ファイルが大きい理由を数える
resume.pthで、重みの要素数(lm_head.weightは埋め込みと共有なので除く)と、オプティマイザーの状態のテンソル(次元が1以上のもの)の要素数、その倍率を、model_numel・optimizer_numel・ratioとして書いてください(出力先: /root/mm/ckpt/optim.json)。
optimizer.state_dict()['state']は、パラメーターごとにstep(0次元)・exp_avg・exp_avg_sqを持っています。倍率が2なら、再開ファイルは、重みファイルの約3倍です。大きなモデルのチェックポイントが重い理由です。
何をなぜ保存するのかを残す
## 무엇을 저장하나 ## 빠뜨리면 ## 반정밀도の3つのセクションを書き、ステップ7のratioとステップ6のfp16_bytesを数字で入れてください(出力先: /root/mm/ckpt/report.md)。韓国語の見出しは、順に「何を保存するのか」「抜いたら」「半精度」という意味です。
ステップ4・5で開いた差の大きさも一緒に書くとよいです。半精度のセクションには、「推論用の重み」と「再開用の重み」をどう分けるかを、1行で書いてみてください。