Articulo de referencia

Google JAX

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 auto...

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]

  1. grad: diferenciación automática
  2. jit: compilación
  3. vmap: auto-vectorización
  4. 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.

Vídeo ilustrativo de la suma 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

  • 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

  1. ^ "jax/AUTHORS at main · jax-ml/jax" . Consultado el 21 de diciembre de 2024 .
  2. ^ 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
  3. ^ 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 )
  4. ^ "Usando JAX para acelerar nuestra investigación". www.deepmind.com . Archivado desde el original el 2022-06-18 . Consultado el 2022-06-18 .
  5. ^ 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 .
  6. ^ "¿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 .
Obtenido de "https://es.wikipedia.org/w/index.php?title=Google_JAX&oldid=1264258322"