HomebrewML/HeavyBall

Efficient optimizers

What it solves

HeavyBallは、PyTorch向けの高性能で構成可能なオプティマイザライブラリを提供します。標準的なオプティマイザの実装における非効率性を、Tritonカーネルに融合(fuse)されるコンパイル済みの構成要素を使用することで解決し、メモリトラフィックを削減し、実行速度を向上させます。また、通常は実装や組み合わせが困難な高度な最適化手法(2次手法や状態圧縮など)の統合を簡素化します。

How it works

このライブラリは、100以上のコンパイル済み関数から構築されており、それらは変換(transforms)の連鎖として組み立てられます。torch.compile(fullgraph=True)を使用することで、これらの関数は最小限のカーネルに融合され、メモリへの読み書きの回数を大幅に削減します。例えば、標準的なAdamの更新は、14回の読み込みと9回の書き込みから、4回の読み込みと3回の書き込みに削減できます。

このライブラリは、1次手法(AdamW, SGD)、直交手法(Muon)、Shampooベース(SOAP)、およびKronecker分解ベース(PSGD)を含む、幅広いオプティマイザのタイプをサポートしています。また、異なるパラメータグループに異なるオプティマイザを適用するための「SplitOpt」も備えています。分散トレーニングにおいては、FSDPを使用する際に2次手法のための再分割(repartitioning)を自動的に処理します。

Who it’s for

より高速なオプティマイザのステップ、オプティマイザ状態のメモリオーバーヘッドの低減、またはPyTorchエコシステム内での多様な高度な2次手法や直交最適化アルゴリズムへのアクセスを必要とする機械学習エンジニアおよび研究者。

Highlights

  • Extensive Optimizer Suite: AdamW, SGD, RMSpropのAPI互換の置き換えに加え、Muon, SOAP, LATHER, ADOPTを含む広範なオプティマイザスイート。
  • Composable Features: MARS分散低減、慎重な更新(cautious updates)、およびPaLMスタイルのbeta2スケジューリングのための、連鎖可能なフラグ。
  • ECC State Compression: 精度を犠牲にすることなく、スペースを節約するためにオプティマイザ状態のメモリ使用量を削減します(例:bf16 + int8 correction)。
  • High Performance: torch.compileを介して操作をTritonカーネルに融合し、ステップのレイテンシを大幅に向上させます。
  • Distributed Support: DDPおよびFSDPとのネイティブな互換性。複雑な2次手法のための自動再分割を含みます。

関連

  • プロジェクト
  • プロジェクト
  • プロジェクト
  • プロジェクト
  • プロジェクト