JAX trabaja de forma natural con arrays y, cuando se necesitan datos agregados, ofrece los pytrees como mecanismo estándar para empaquetar varios arrays en una sola estructura que el sistema aplana y reconstruye de forma transparente. Sin embargo, hay casos en los que esa transparencia resulta contraproducente: cuando se quiere modelar ciertos datos como un tipo nuevo con identidad propia, con invariantes internos, con tangentes diferenciados respecto a la estructura principal y con semántica propia de batching y particionado. Para esos casos, JAX introduce los tipos hijax (o "hi types").
Un tipo hijax se define como subclase de HiType y se asocia a una clase portadora mediante register_hitype. Sus métodos clave son lo_ty, lower_val y raise_val, que describen cómo el tipo y sus valores se reducen a arrays ordinarios ("lojax") y cómo se reconstruyen a partir de ellos. La creación y el consumo de instancias se canalizan exclusivamente a través de primitivas hijax, normalmente subclases de VJPHiPrimitive, cuyos tipos de entrada y salida mencionan el nuevo tipo. De este modo, los invariantes los garantiza el propio sistema de tipos, no la convención del usuario.
El artículo desarrolla el flujo completo con un ejemplo: un tipo QArray que representa un array cuantizado a int8 con una escala float32 por fila. El texto explica cómo declarar el tipo con su campo de particionado (sharding), cómo registrarlo, cómo implementar las primitivas quantize y dequantize, cómo definir reglas de diferenciación automática mediante to_tangent_aval y reglas VJP/JVP, y cómo soportar vmap con MappingSpec, dec_rank, inc_rank y batchrules. También aborda el modo de particionado explícito, registrando NamedSharding en el tipo y propagándolo en lo_ty y en las reglas de tipado de las primitivas. Todo el contenido es experimental y reside en jax.experimental.hijax, por lo que las API pueden cambiar.
