Sobre o projeto
O JAX fornece um sistema para transformações de funções compostíveis de programas Python e NumPy. Ele foi projetado para escalar a computação numérica em aceleradores de hardware, incluindo GPUs e TPUs, utilizando XLA (Accelerated Linear Algebra) para compilação.
As principais capacidades incluem:
- Diferenciação Automática: Usando `jax.grad`, suporta diferenciação de modo reverso e modo direto através de loops Python, ramificações, recursão e closures, permitindo derivadas de ordem superior.
- Compilação Just-In-Time (JIT): Usando `jax.jit`, compila funções puras via XLA para melhorar o desempenho de execução em aceleradores.
- Auto-vetorização: Usando `jax.vmap`, mapeia funções ao longo de eixos de arrays, empurrando loops para operações primitivas, eliminando a necessidade de gerenciamento manual de dimensões de lote.
- Escalonamento e Paralelização: Suporta vários modos de escalonamento, incluindo paralelização automática baseada em compilador, sharding explícito e programação manual por dispositivo com coletivos explícitos.
O JAX é compatível com múltiplas plataformas, incluindo Linux, macOS e Windows, com suporte para CPU, NVIDIA GPU, Google TPU, AMD GPU e suporte experimental para GPUs Apple e Intel.