Profilage dans PyTorch (Partie 3) : L'attention est tout ce qu'il faut profiler
PyTorch
NVIDIA
Cet article profile les mécanismes d'attention dans PyTorch, en comparant l'attention naïve, l'optimisation sur place et les backends de l'attention produit scalaire mise à l'échelle (SDPA). Il révèle que le backend mathématique de SDPA est plus lent que l'attention naïve en raison de la sous-utilisation des Tensor Cores, de la reconstruction du masque et du surcoût du softmax sécurisé.
Le troisième article de la série Profiling in PyTorch profile les mécanismes d'attention. Il commence par un module d'attention naïf (matmul, mise à l'échelle, masquage, softmax, matmul) et sa trace montre 5 noyaux GPU par forward. Remplacer masked_fill hors place par masked_fill_ en place élimine un noyau de copie mémoire, économisant temps et mémoire. Ensuite, l'Attention par Produit Scalaire Pondéré (SDPA) est présentée ; son backend mathématique, bien que plus sûr (gérant les cas NaN) et plus précis (FP32), est 3,7 fois plus lent que l'attention naïve car il lance 20 noyaux GPU, n'utilise pas les Tensor Cores (utilise sgemm au lieu de matmul bf16 Tensor Core), reconstruit le masque causal à chaque appel, et utilise _safe_softmax. L'article note que SDPA sélectionne normalement le backend le plus rapide automatiquement.
Source: Hugging Face blog —
original
