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.