Об этом проекте
Transformer Engine (TE) — это библиотека NVIDIA для ускорения Transformer-моделей на GPU NVIDIA. Её центральная особенность — вычисления с низкой точностью: 8-битная плавающая точка (FP8) на GPU Hopper, Ada и Blackwell, а также форматы MXFP8 и NVFP4 на Blackwell, предназначенные для повышения производительности и снижения использования памяти как при обучении, так и при инференсе. Она также поддерживает оптимизации для FP16 и BF16 на архитектурах Ampere и более новых.
Библиотека предоставляет высокооптимизированные строительные блоки для распространённых архитектур Transformer и API, подобный автоматической смешанной точности, который можно использовать с кодом, специфичным для фреймворка. Включён фреймворк-агностичный C++ API, чтобы другие библиотеки глубокого обучения могли добавить поддержку FP8 для Transformer. Модули TE внутренне поддерживают коэффициенты масштабирования и связанные значения, необходимые для обучения с FP8, что упрощает рабочие процессы смешанной точности для пользователей.
Среди ключевых возможностей, перечисленных в проекте: простые в использовании модули для построения слоёв Transformer с поддержкой FP8, оптимизации fused-ядер, поддержка FP8 на Hopper, Ada и Blackwell, поддержка MXFP8 и NVFP4 на Blackwell, а также оптимизации для FP16/BF16 на Ampere и более новых.
Примеры использования приведены для PyTorch и JAX/Flax. В PyTorch пользователи импортируют transformer_engine.pytorch, создают рецепт, например DelayedScaling с выбранным форматом FP8, и оборачивают прямой проход в te.autocast. В JAX/Flax показано аналогичное использование autocast с модулями te_flax и рецептом. Для более полного руководства приведена ссылка на getting-started guide.
Варианты установки включают Docker-контейнеры NGC (рекомендуется), pip-пакеты с extras для PyTorch и/или JAX, пакеты conda-forge для интеграции с PyTorch и сборку из исходников. Системные требования упоминают оборудование Blackwell, Hopper, Grace Hopper/Blackwell, Ada и Ampere; Linux как официальную ОС с ограниченной поддержкой WSL2; CUDA 12.1+ (12.8+ для Blackwell); cuDNN 9.12+; GCC 9+ или Clang 10+ с C++17; рекомендуется Python 3.12. Функции FP8 требуют compute capability 8.9 или выше. Документированы переменные окружения, связанные со сборкой, такие как CUDA_PATH, CUDNN_PATH, CXX, NVTE_FRAMEWORK, MAX_JOBS и NVTE_CUDA_ARCHS.
README описывает поддержку FlashAttention-2 и FlashAttention-3 в PyTorch, при этом FlashAttention-3 имеет приоритет, когда присутствуют обе, и отмечает, что компиляция FlashAttention-2 может быть ресурсоёмкой по памяти. Раздел устранения неполадок рассматривает ошибки импорта, связанные с совместимостью ABI, отсутствующие заголовки или библиотеки, проблемы с ресурсами при сборке, подробное логирование сборки, проблемы UV/виртуальных окружений, включая сбои загрузки подбиблиотек cuDNN, и ошибки регистрации JAX FFI.
Документировано критическое изменение в v1.7: определение padding-маски в PyTorch изменилось так, что True теперь означает маскирование позиции, а не её включение, что унифицирует семантику масок между фреймворками. Замечания о сходимости указывают, что FP8 и MXFP8 не показали значительных отличий от кривых потерь обучения BF16 в протестированных конфигурациях, с валидацией на последующих задачах LLM, и перечисляют такие модели, как MPT-1.3B, Llama2-7B, LLM-8B, MPT-13B, MoE-16B и Llama2-70B, в рамках фреймворков, включая Mosaic Composer, Alibaba Pai и Megatron Core.
Среди перечисленных интеграций: 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 и Hugging Face Nanotron. Проект лицензирован под Apache-2.0 и приветствует вклад через свой CONTRIBUTING guide.
Comments
0 people shared their preference · Deer Point appears after 10 participants
Sign in to join the discussion.