hijax: define nuevos tipos en JAX con identidad propia, más allá de los pytrees
Defining new Jax types with hijax
hijax, la API experimental de JAX, permite definir tipos personalizados que se comportan como ciudadanos de primera clase en el ecosistema JAX. A diferencia de los pytrees, que se aplanan en arrays en cada frontera, los tipos hijax mantienen su identidad en los jaxprs, pueden tener invariantes internos, tipos tangentes personalizados, reglas de batching propias y sharding en el tipo. El documento muestra cómo implementar un tipo de array cuantizado (q8) con primitivas hijax, incluyendo reglas de VJP y vmap, y destaca la importancia de usar primitivas fuera de la implementación para preservar invariantes.
A veces la transparencia es exactamente lo que no quieres.