JAX es un marco de aprendizaje automático para transformar funciones numéricas. [2] [3] [4] Se describe como la combinación de una versión modificada de autograd (obtención automática de la función de gradiente a través de la diferenciación de una función) y XLA (Álgebra lineal acelerada) de OpenXLA. Está diseñado para seguir la estructura y el flujo de trabajo de NumPy lo más fielmente posible y funciona con varios marcos existentes como TensorFlow y PyTorch . [5] [6] Las funciones principales de JAX son: [2]
- grad: diferenciación automática
- jit: compilación
- vmap: auto-vectorización
- pmap: Programación de un solo programa y múltiples datos (SPMD)
graduado
El siguiente código demuestra la diferenciación automática de la función grad .
# importaciones
de jax import grad
importar jax.numpy como jnp
# definir la función logística
def logística ( x ):
devuelve jnp . exp ( x ) / ( jnp . exp ( x ) + 1 )
# obtener la función gradiente de la función logística
grad_logistic = grad ( logístico )
# evaluar el gradiente de la función logística en x = 1
salida_grad_log = logística_grad ( 1.0 )
imprimir ( grad_log_out )
La línea final debería mostrar:ː
0,19661194
escurridizo
El siguiente código demuestra la optimización de la función jit a través de la fusión.
# importaciones
Desde jax importar jit
importar jax.numpy como jnp
# define la función cubo
def cubo ( x ):
devolver x * x * x
# generar datos
x = jnp . unos (( 10000 , 10000 ))
# crea la versión jit de la función cubo
jit_cube = jit ( cubo )
# aplicar las funciones cube y jit_cube a los mismos datos para comparar la velocidad
cubo ( x )
cubo jit ( x )
El tiempo de cálculo de jit_cubela línea n.° 17 debería ser notablemente más corto que el de cubela línea n.° 16. Si se aumentan los valores de la línea n.° 7, la diferencia se acentuará aún más.
Mapa virtual
El siguiente código demuestra la vectorización de la función vmap .
# importaciones
desde jax import vmap parcial
importar jax.numpy como jnp
# definir función
def grads ( self , entradas ):
in_grad_partial = jax . partial ( self . _net_grads , self . _net_params )
grad_vmap = jax . vmap ( in_grad_partial )
rich_grads = grad_vmap ( entradas )
flat_grads = np . asarray ( self . _flatten_batch ( rich_grads ))
afirmar flat_grads . ndim == 2 y flat_grads . shape [ 0 ] == entradas . shape [ 0 ]
Devuelve flat_grads
El GIF a la derecha de esta sección ilustra la noción de adición vectorizada.

mapa p
El siguiente código demuestra la paralelización de la función pmap para la multiplicación de matrices.
# importar pmap y random desde JAX; importar JAX NumPy
desde jax import pmap , aleatorio
importar jax.numpy como jnp
# generar 2 matrices aleatorias de dimensiones 5000 x 6000, una por dispositivo
claves_aleatorias = random.split ( random.PRNGKey ( 0 ) , 2 )
matrices = pmap ( clave lambda : aleatorio.normal ( clave , ( 5000,6000 ) ) ) ( claves_aleatorias )
# sin transferencia de datos, en paralelo, realice una multiplicación de matriz local en cada CPU/GPU
salidas = pmap ( lambda x : jnp . dot ( x , x . T )) ( matrices )
# sin transferencia de datos, en paralelo, obtenga la media de ambas matrices en cada CPU/GPU por separado
media = pmap ( jnp.media ) ( salidas )
imprimir ( significa )
La línea final debe imprimir los valoresː
[1.1566595 1.1805978]
Véase también
Enlaces externos
- Documentaciónː jax.readthedocs.io
- Guía de inicio rápido de Colab ( Jupyter /iPython)ː colab.research.google.com/github/google/jax/blob/main/docs/notebooks/quickstart.ipynb
- XLAː de TensorFlow www.tensorflow.org/xla (Álgebra lineal acelerada)
- Canal de YouTube TensorFlow "Introducción a JAX: Aceleración de la investigación en aprendizaje automático": www.youtube.com/watch?v=WdTeDXsOSj4
- Documento original: mlsys.org/Conferences/doc/2018/146.pdf
Referencias
- ^ "jax/AUTHORS at main · jax-ml/jax" . Consultado el 21 de diciembre de 2024 .
- ^ ab Bradbury, James; Frostig, Roy; Hawkins, Peter; Johnson, Matthew James; Leary, Chris; MacLaurin, Dougal; Necula, George; Paszke, Adam; Vanderplas, Jake; Wanderman-Milne, Skye; Zhang, Qiao (18 de junio de 2022), "JAX: Autograd y XLA", Astrophysics Source Code Library , Google, Bibcode :2021ascl.soft11002B, archivado desde el original el 18 de junio de 2022 , consultado el 18 de junio de 2022
- ^ Frostig, Roy; Johnson, Matthew James; Leary, Chris (2 de febrero de 2018). "Compiling machine learning programs via high-level tracing" (PDF) . MLsys : 1– 3. Archivado (PDF) desde el original el 21 de junio de 2022.
{{cite journal}}: Mantenimiento CS1: fecha y año ( enlace ) - ^ "Usando JAX para acelerar nuestra investigación". www.deepmind.com . Archivado desde el original el 2022-06-18 . Consultado el 2022-06-18 .
- ^ Lynley, Matthew. "Google está reemplazando silenciosamente la columna vertebral de su estrategia de productos de inteligencia artificial después de que su último gran impulso por el dominio se viera eclipsado por Meta". Business Insider . Archivado desde el original el 2022-06-21 . Consultado el 2022-06-21 .
- ^ "¿Por qué JAX de Google es tan popular?". Revista Analytics India . 25 de abril de 2022. Archivado desde el original el 18 de junio de 2022. Consultado el 18 de junio de 2022 .