noahgolmant/pytorch-hessian-eigenthings

Efficient PyTorch Hessian eigendecomposition tools!

What it solves

大規模なニューラルネットワークにおいて、フルヘシアン行列(損失関数の2次導関数)を計算することは、メモリ要件がパラメータ数に対して二次関数的に増加するため、計算的に不可能です。このライブラリは、フル行列を明示的に構築することなく、モデルの損失ランドスケープの曲率を分析するためのスケーラブルな方法を提供します。

How it works

このプロジェクトは、線形メモリのみを必要とするヘシアン・ベクトル積(HVP)を使用しています。HVPを反復アルゴリズムと組み合わせることで、メモリのボトルネックなしに、ヘシアンやその他の曲率行列(Generalized Gauss-Newtonやempirical Fisherなど)の特定の特性を計算できます。

主要なアルゴリズムには以下が含まれます:

  • Lanczos and stochastic power iteration は、上位の固有値と固有ベクトルを見つけるためのものです。
  • Hutch++ は、行列のトレースを推定するためのものです。
  • Stochastic Lanczos Quadrature は、スペクトル密度を推定するためのものです。

大規模言語モデル(LLM)向けには、クロスエントロピーのヘシアン・ベクトル積を高速化するために、最適化されたカーネル(Tritonまたはtorch.compileを使用)が含まれています。

Who it’s for

ニューラルネットワークの最適化、汎化分析、および損失ランドスケープの幾何学(例:「flat minima」の分析)を研究している研究者や開発者。

Highlights

  • Scalable Analysis: HuggingFaceやTransformerLensのモデルを含む、実世界のモデルの固有値分解とスペクトル密度を計算します。
  • Flexible Operators: Hessian、Generalized Gauss-Newton (GGN)、およびEmpirical Fisherオペレータをサポートします。
  • Memory Efficient: 二次メモリではなく、HVPを介した線形メモリを使用します。
  • Performance Optimizations: ピークメモリを削減し、速度を向上させるために、LM規模の作業向けに融合カーネル(fused kernels)を搭載しています。
  • Parameter Filtering: 名前ベースのフィルタを使用して、特定のパラメータサブセット(例:transformer内の特定のブロック)の分析を可能にします。

関連

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