Sobre o projeto
O Transformer Engine (TE) é uma biblioteca da NVIDIA para acelerar modelos Transformer em GPUs NVIDIA. Seu recurso central é a computação de baixa precisão: ponto flutuante de 8 bits (FP8) em GPUs Hopper, Ada e Blackwell, além dos formatos MXFP8 e NVFP4 em Blackwell, com o objetivo de melhorar o desempenho e reduzir o uso de memória tanto no treinamento quanto na inferência. Também oferece suporte a otimizações em FP16 e BF16 nas arquiteturas Ampere e posteriores.
A biblioteca fornece blocos de construção altamente otimizados para arquiteturas Transformer comuns e uma API automática semelhante à de precisão mista, que pode ser usada com código específico de cada framework. Uma API C++ independente de framework está incluída para que outras bibliotecas de aprendizado profundo possam adicionar suporte a FP8 para Transformers. Os módulos do TE mantêm internamente fatores de escala e valores relacionados necessários para o treinamento em FP8, o que simplifica os fluxos de trabalho de precisão mista para os usuários.
Os destaques listados pelo projeto incluem módulos fáceis de usar para construir camadas Transformer com suporte a FP8, otimizações de kernels fundidos, suporte a FP8 em Hopper, Ada e Blackwell, suporte a MXFP8 e NVFP4 em Blackwell, e otimizações em FP16/BF16 em Ampere e posteriores.
Exemplos de uso são fornecidos para PyTorch e JAX/Flax. No PyTorch, os usuários importam transformer_engine.pytorch, criam uma receita como DelayedScaling com um formato FP8 escolhido e envolvem o forward pass em te.autocast. No JAX/Flax, um uso semelhante de autocast é mostrado com módulos te_flax e uma receita. Um guia de introdução está vinculado para um tutorial mais completo.
As opções de instalação incluem contêineres Docker NGC (recomendado), pacotes pip com extras para PyTorch e/ou JAX, pacotes conda-forge para a integração com PyTorch e compilações a partir do código-fonte. Os requisitos de sistema mencionam hardware Blackwell, Hopper, Grace Hopper/Blackwell, Ada e Ampere; Linux como sistema operacional oficial, com suporte limitado a WSL2; CUDA 12.1+ (12.8+ para Blackwell); cuDNN 9.12+; GCC 9+ ou Clang 10+ com C++17; e Python 3.12 recomendado. Os recursos de FP8 exigem capacidade de computação 8.9 ou superior. Variáveis de ambiente relacionadas à compilação, como CUDA_PATH, CUDNN_PATH, CXX, NVTE_FRAMEWORK, MAX_JOBS e NVTE_CUDA_ARCHS, estão documentadas.
O README aborda o suporte a FlashAttention-2 e FlashAttention-3 no PyTorch, com prioridade para FlashAttention-3 quando ambos estão presentes, e observa que a compilação do FlashAttention-2 pode consumir muita memória. Uma seção de solução de problemas aborda erros de importação por compatibilidade de ABI, cabeçalhos ou bibliotecas ausentes, problemas de recursos de compilação, registro detalhado de compilação, problemas com UV/ambiente virtual, incluindo falhas de carregamento de subbibliotecas do cuDNN, e erros de registro de FFI do JAX.
Uma mudança significativa está documentada para a v1.7: a definição da máscara de preenchimento do PyTorch mudou, de modo que True agora significa mascarar uma posição em vez de incluí-la, unificando a semântica de máscaras entre os frameworks. As notas de convergência afirmam que FP8 e MXFP8 não mostraram diferença significativa em relação às curvas de perda de treinamento em BF16 nas configurações testadas, com validação em tarefas downstream de LLM, e listam modelos como MPT-1.3B, Llama2-7B, LLM-8B, MPT-13B, MoE-16B e Llama2-70B em frameworks como Mosaic Composer, Alibaba Pai e Megatron Core.
As integrações listadas incluem 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 e Hugging Face Nanotron. O projeto é licenciado sob Apache-2.0 e aceita contribuições por meio de seu guia CONTRIBUTING.
Comments
0 people shared their preference · Deer Point appears after 10 participants
Sign in to join the discussion.