Skip to content

勾配クリッピング

Gradient Clipping

Gradient ClippingGlobal norm threshold 1.0: a single bad batch rescaled, training continues uninterrupted01020304050607080Training step0123456Gradient normclip = 1.0LossClipping modesGlobal norm clipscale all param gradsif ||g|| > tPer-tensor clipeach layer separatelypreserves ratiosValue clipclamp each elementrarer; cruderClipped norm / lossRaw norm / lossGlobal norm clip (threshold 1.0) is the default in PyTorch, JAX, and Hugging Face Trainer for transformer pretraining

勾配クリッピングは、勾配ベクトルのグローバルL2ノルムが閾値tを超えた場合に、勾配を再スケーリングする手法である。演算は次のように定義される。||g|| > t のとき、g = g × (t / ||g||)。これにより勾配の向きを保持しつつ、その大きさをちょうどtに制限する。この点で、勾配の向きを歪める要素ごとの値クリッピングとは異なる。実際には閾値を0.5から5.0の間に設定する。GPT-2とGPT-3はt = 1.0、LLaMAは1.0、PaLMは1.0を使用しており、1.0がTransformerの事前学習における事実上の標準となっている。損失関数が極端な入力に遭遇した際に生じる単一の異常バッチは、通常の10〜100倍の勾配ノルムを引き起こすことがある。クリッピングなしでは重みが高損失領域に大きくステップし、回復不能な損失スパイクを招く恐れがある。この図では、ステップ34において勾配ノルムが6.1に急上昇している(通常値は約0.4)。クリッピングなしでは訓練損失が発散し、1.0でクリッピングするとノルムが再スケーリングされ、訓練はスムーズに継続される。

English version