Sobre el proyecto

JAX proporciona un sistema para transformaciones de funciones componibles de programas de Python y NumPy. Está diseñado para escalar el cómputo numérico a través de aceleradores de hardware, incluyendo GPUs y TPUs, utilizando XLA (Accelerated Linear Algebra) para la compilación. Las capacidades clave incluyen: - Diferenciación Automática: Usando `jax.grad`, soporta la diferenciación en modo inverso y modo directo a través de bucles de Python, ramificaciones, recursión y clausuras, permitiendo derivadas de orden superior. - Compilación Just-In-Time (JIT): Usando `jax.jit`, compila funciones puras a través de XLA para mejorar el rendimiento de ejecución en aceleradores. - Auto-vectorización: Usando `jax.vmap`, mapea funciones a lo largo de los ejes de los arreglos desplazando los bucles hacia las operaciones primitivas, eliminando la necesidad de gestionar manualmente las dimensiones de lote. - Escalado y Paralelización: Soporta varios modos de escalado, incluyendo la paralelización automática basada en el compilador, el sharding explícito y la programación manual por dispositivo con colectivos explícitos. JAX es compatible con múltiples plataformas, incluyendo Linux, macOS y Windows, con soporte para CPU, NVIDIA GPU, Google TPU, AMD GPU y soporte experimental para GPUs de Apple e Intel.