Sobre el proyecto

Transformer Engine (TE) es una biblioteca de NVIDIA para acelerar modelos Transformer en GPU NVIDIA. Su característica central es la computación de baja precisión: punto flotante de 8 bits (FP8) en GPU Hopper, Ada y Blackwell, además de los formatos MXFP8 y NVFP4 en Blackwell, destinados a mejorar el rendimiento y reducir el uso de memoria tanto en entrenamiento como en inferencia. También admite optimizaciones en FP16 y BF16 en arquitecturas Ampere y posteriores. La biblioteca proporciona bloques de construcción altamente optimizados para arquitecturas Transformer comunes y una API automática similar a la precisión mixta que puede usarse con código específico de cada framework. Se incluye una API C++ independiente del framework para que otras bibliotecas de aprendizaje profundo puedan añadir soporte FP8 para Transformers. Los módulos de TE mantienen internamente los factores de escala y los valores relacionados necesarios para el entrenamiento FP8, lo que simplifica los flujos de trabajo de precisión mixta para los usuarios. Entre los aspectos destacados que enumera el proyecto se incluyen módulos fáciles de usar para construir capas Transformer con soporte FP8, optimizaciones de kernels fusionados, soporte FP8 en Hopper, Ada y Blackwell, soporte MXFP8 y NVFP4 en Blackwell, y optimizaciones en FP16/BF16 en Ampere y posteriores. Se ofrecen ejemplos de uso para PyTorch y JAX/Flax. En PyTorch, los usuarios importan transformer_engine.pytorch, crean una receta como DelayedScaling con un formato FP8 elegido y envuelven el pase forward en te.autocast. En JAX/Flax, se muestra un uso similar de autocast con módulos te_flax y una receta. Se enlaza una guía de introducción para un tutorial más completo. Las opciones de instalación incluyen contenedores Docker de NGC (recomendado), paquetes pip con extras para PyTorch y/o JAX, paquetes de conda-forge para la integración con PyTorch y compilaciones desde el código fuente. Los requisitos del sistema mencionan hardware Blackwell, Hopper, Grace Hopper/Blackwell, Ada y Ampere; Linux como sistema operativo oficial con soporte limitado para WSL2; CUDA 12.1+ (12.8+ para Blackwell); cuDNN 9.12+; GCC 9+ o Clang 10+ con C++17; y Python 3.12 recomendado. Las funciones FP8 requieren capacidad de cómputo 8.9 o superior. Se documentan variables de entorno relacionadas con la compilación, como CUDA_PATH, CUDNN_PATH, CXX, NVTE_FRAMEWORK, MAX_JOBS y NVTE_CUDA_ARCHS. El README cubre el soporte de FlashAttention-2 y FlashAttention-3 en PyTorch, con prioridad para FlashAttention-3 cuando ambos están presentes, y señala que la compilación de FlashAttention-2 puede consumir mucha memoria. Una sección de solución de problemas aborda errores de importación por compatibilidad ABI, encabezados o bibliotecas faltantes, problemas de recursos de compilación, registro detallado de compilación, problemas con UV/entornos virtuales, incluidos fallos de carga de sublibrerías de cuDNN, y errores de registro de JAX FFI. Se documenta un cambio incompatible para la v1.7: la definición de la máscara de padding de PyTorch cambió, de modo que True ahora significa enmascarar una posición en lugar de incluirla, unificando la semántica de las máscaras entre frameworks. Las notas de convergencia indican que FP8 y MXFP8 no mostraron diferencias significativas respecto a las curvas de pérdida de entrenamiento con BF16 en las configuraciones probadas, con validación en tareas posteriores de LLM, y enumeran modelos como MPT-1.3B, Llama2-7B, LLM-8B, MPT-13B, MoE-16B y Llama2-70B en frameworks como Mosaic Composer, Alibaba Pai y Megatron Core. Las integraciones enumeradas incluyen 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 y Hugging Face Nanotron. El proyecto tiene licencia Apache-2.0 y acepta contribuciones a través de su guía CONTRIBUTING.