Skip to content

Écrêtage de gradient

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

L'écrêtage de gradient redimensionne le vecteur gradient lorsque sa norme L2 globale dépasse un seuil t. L'opération est la suivante : si ||g|| > t, alors g = g * (t / ||g||). Cela préserve la direction du gradient tout en bornant sa magnitude exactement à t, contrairement à l'écrêtage élément par élément qui distord la direction. En pratique, le seuil est fixé entre 0,5 et 5,0 : GPT-2 et GPT-3 utilisent t = 1,0, LLaMA utilise 1,0, et PaLM utilise 1,0, faisant de 1,0 la valeur de référence de facto pour le préentraînement des transformeurs. Un seul lot anormal, où la fonction de perte rencontre une entrée extrême, peut produire une norme de gradient 10 à 100 fois supérieure à la valeur habituelle ; sans écrêtage, cela projette les poids dans une grande enjambée vers une région à perte élevée, provoquant un pic de perte dont la récupération n'est pas garantie. Le diagramme montre un pic de norme de gradient à 6,1 à l'étape 34 (contre une valeur typique d'environ 0,4) : sans écrêtage, la perte d'entraînement diverge ; avec un écrêtage à 1,0, la norme est redimensionnée et l'entraînement se poursuit sans à-coups.

English version