À propos du projet
Transformer Engine (TE) est une bibliothèque NVIDIA pour accélérer les modèles Transformer sur les GPU NVIDIA. Sa caractéristique centrale est le calcul en basse précision : virgule flottante 8 bits (FP8) sur les GPU Hopper, Ada et Blackwell, plus les formats MXFP8 et NVFP4 sur Blackwell, destinés à améliorer les performances et à réduire l'utilisation de la mémoire en entraînement comme en inférence. Elle prend également en charge des optimisations sur FP16 et BF16 sur les architectures Ampere et ultérieures.
La bibliothèque fournit des blocs de construction hautement optimisés pour les architectures Transformer courantes et une API de type précision mixte automatique utilisable avec du code spécifique à un framework. Une API C++ indépendante du framework est incluse afin que d'autres bibliothèques d'apprentissage profond puissent ajouter la prise en charge FP8 pour les Transformers. Les modules TE maintiennent en interne les facteurs d'échelle et les valeurs associées nécessaires à l'entraînement FP8, ce qui simplifie les flux de travail en précision mixte pour les utilisateurs.
Les points forts listés par le projet incluent des modules faciles à utiliser pour construire des couches Transformer avec prise en charge FP8, des optimisations par noyaux fusionnés, la prise en charge FP8 sur Hopper, Ada et Blackwell, la prise en charge MXFP8 et NVFP4 sur Blackwell, et des optimisations sur FP16/BF16 sur Ampere et ultérieurs.
Des exemples d'utilisation sont donnés pour PyTorch et JAX/Flax. En PyTorch, les utilisateurs importent transformer_engine.pytorch, créent une recette telle que DelayedScaling avec un format FP8 choisi, et enveloppent la passe avant dans te.autocast. En JAX/Flax, un usage similaire d'autocast est montré avec les modules te_flax et une recette. Un guide de démarrage est lié pour un tutoriel plus complet.
Les options d'installation incluent les conteneurs Docker NGC (recommandés), les paquets pip avec extras pour PyTorch et/ou JAX, les paquets conda-forge pour l'intégration PyTorch, et les compilations depuis les sources. Les prérequis système mentionnent le matériel Blackwell, Hopper, Grace Hopper/Blackwell, Ada et Ampere ; Linux comme système d'exploitation officiel avec prise en charge limitée de WSL2 ; CUDA 12.1+ (12.8+ pour Blackwell) ; cuDNN 9.12+ ; GCC 9+ ou Clang 10+ avec C++17 ; et Python 3.12 recommandé. Les fonctionnalités FP8 nécessitent une capacité de calcul 8.9 ou supérieure. Les variables d'environnement liées à la compilation telles que CUDA_PATH, CUDNN_PATH, CXX, NVTE_FRAMEWORK, MAX_JOBS et NVTE_CUDA_ARCHS sont documentées.
Le README couvre la prise en charge de FlashAttention-2 et FlashAttention-3 en PyTorch, avec FlashAttention-3 priorisé lorsque les deux sont présents, et note que la compilation de FlashAttention-2 peut être gourmande en mémoire. Une section de dépannage traite des erreurs d'import de compatibilité ABI, des en-têtes ou bibliothèques manquants, des problèmes de ressources de compilation, de la journalisation verbeuse de compilation, des problèmes UV/environnement virtuel incluant les échecs de chargement des sous-bibliothèques cuDNN, et des erreurs d'enregistrement JAX FFI.
Un changement incompatible est documenté pour la v1.7 : la définition du masque de remplissage PyTorch a changé, de sorte que True signifie désormais masquer une position plutôt que l'inclure, unifiant la sémantique des masques entre les frameworks. Les notes de convergence indiquent que FP8 et MXFP8 n'ont montré aucune différence significative par rapport aux courbes de perte d'entraînement BF16 dans les configurations testées, avec validation sur des tâches LLM en aval, et listent des modèles tels que MPT-1.3B, Llama2-7B, LLM-8B, MPT-13B, MoE-16B et Llama2-70B à travers des frameworks incluant Mosaic Composer, Alibaba Pai et Megatron Core.
Les intégrations listées incluent DeepSpeed, Hugging Face Accelerate, Lightning, MosaicML Composer, NVIDIA JAX Toolbox, NVIDIA Megatron-LM, NVIDIA NeMo Megatron Bridge, Amazon SageMaker Model Parallel Library, Levanter, GPT-NeoX et Hugging Face Nanotron. Le projet est sous licence Apache-2.0 et accueille les contributions via son guide CONTRIBUTING.
Comments
0 people shared their preference · Deer Point appears after 10 participants
Sign in to join the discussion.