如何用 hijax 定义新的 JAX 类型,让它拥有自己的切线类型和批量规则
【免费下载链接】jaxComposable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax
在 JAX 中,转换函数处理的内置数据类型是数组:函数进出都是数组,tracing 机制看到的每个中间值都是f32[3,4]这样的数组类型。当你有一类带内部不变量的聚合数据(比如量化数组:整数载荷加每行共享的浮点 scale),把它注册成 pytree 意味着放弃三件事:jaxpr 中它会显示成两个互不相关的数组值、无法指定一个与值结构不同的切线类型、vmap的批量语义也只能按普通 pytree 处理。
JAX 的 hijax 扩展(hi types,即 “hi 类型”)就是为此设计的:你子类化HiType定义类型、用register_hitype把值类关联到该类型、写HiPrim原语来生产/消费该类型的值,然后为自动微分和vmap分别提供切线类型和批量规则。本文按 Defining new JAX types with hijax 文档的完整示例,走一遍“定义一个int8量化数组类型,让它拥有自己的切线类型和批量规则”的全过程,并在每一步给出文档中的验证方式。
前提说明:hijax 整体仍是实验特性,导入来自jax.experimental.hijax,API 会持续演进;文档建议先熟悉 hijax 原语的基本用法,见 自定义导数规则文档。
准备条件:导入与值类
文档示例定义一个按行量化的数组:int8的qvalue加上每行一个f32的scale。先定义值类(文档的第一个代码单元):
import os os.environ["XLA_FLAGS"] = '--xla_force_host_platform_device_count=8' # (8 CPU devices, for the sharding sections at the end) from dataclasses import dataclass import jax import jax.numpy as jnp @dataclass(frozen=True) class QArray: qvalue: jax.Array # int8[*leading, n] scale: jax.Array # f32[*leading]注意:文档开头设置XLA_FLAGS强制 8 个 CPU 设备,注释明确这是给文档末尾 sharding 章节用的。如果你只走本文的主路径(定义类型、切线类型、批量规则),这一行不是必需的;如果后续要做显式 sharding 示例,再保留它。
定义类型:HiType子类加register_hitype
hijax 类型是HiType的子类,必须实现的核心很小:
lo_ty:说出这个类型由哪些 lojax(数组)类型组成;lower_val/raise_val:把值和这个数组列表互相转换;- 类型本身必须可哈希且可按相等性比较(frozen dataclass 同时满足两者)。
这类似 pytree 的 flatten/unflatten 接口,但处在类型层面:只给定类型,JAX 就能算出 lower 后的类型,不需要拿到具体值。
文档的完整类型定义如下(sharding字段服务于显式 sharding 模式,没有 mesh 时可以忽略,详见后文“边界”一节):
from jax.experimental.hijax import HiType, ShapedArray, register_hitype from jax.sharding import NamedSharding @dataclass(frozen=True) class QArrayTy(HiType): shape: tuple[int, ...] sharding: NamedSharding # qvalue's sharding; scale's is derived from it # lowering: which array types make up this type, and how values convert def lo_ty(self): scale_sharding = self.sharding.update(spec=jax.P(*self.sharding.spec[:-1])) return [ShapedArray(self.shape, jnp.dtype('int8'), sharding=self.sharding), ShapedArray(self.shape[:-1], jnp.dtype('float32'), sharding=scale_sharding)] def lower_val(self, q): return [q.qvalue, q.scale] def raise_val(self, qvalue, scale): return QArray(qvalue, scale) # autodiff: tangents of quantized arrays are plain float arrays (see below) def to_tangent_aval(self): return ShapedArray(self.shape, jnp.dtype('float32'), sharding=self.sharding) # printing, e.g. in jaxprs def str_short(self, short_dtypes=False, mesh_axis_types=False): dims = [str(d) if p is None else f'{d}@{p}' for d, p in zip(self.shape, self.sharding.spec)] return f'q8[{",".join(dims)}]' __repr__ = str_short register_hitype(QArray, lambda q: QArrayTy(q.qvalue.shape, jax.typeof(q.qvalue).sharding))几个关键点,均来自文档正文:
register_hitype把值类与类型关联起来,第二个参数负责从任意值算出它的类型(类似jax.typeof把数组映射到ShapedArray)。注册之后,jax.typeof就能作用于QArray,JAX 的转换也能在任何期望值的地方接受它们。to_tangent_aval是切线类型的声明:量化数组的切线就是普通f32数组。这是 pytree 表达不了的选择——pytree 的切线类型只能是其叶子切线类型的 pytree,而int8的qvalue的切线只能是只能承载平凡载荷的float0数组。str_short只影响打印(例如 jaxpr 中显示为q8[2,3]),与语义无关。
定义原语:值只能由原语生产和消费
用 pytree 时用户可以随意构造和解构值;用 hijax 类型时,值只能由声明类型中提及该类型的 hijax 原语生产和消费。不变量正是在这里被强制的:只要每个原语都保持它,它就永远成立。
示例的两个原语是quantize和dequantize,用 HiPrim API 编写。每个原语在__init__中声明输入输出类型、在expand中给出实现,并(为自动微分预做)携带 straight-through-estimator VJP 规则:
from jax.experimental.hijax import HiPrim class Quantize(HiPrim): def __init__(self, x_aval): if x_aval.dtype != jnp.dtype('float32'): raise TypeError(x_aval.dtype) self.in_avals = (x_aval,) self.out_aval = QArrayTy(x_aval.shape, x_aval.sharding) self.params = {} super().__init__() def expand(self, x): scale = jnp.max(jnp.abs(x), axis=-1) / 127. qvalue = jnp.round(x / scale[..., None]).astype(jnp.int8) return QArray(qvalue, scale) # straight-through estimator: differentiate as if it's the identity def vjp_fwd(self, nzs_in, x): return self(x), None def vjp_bwd_retval(self, _res, g): return (g,) class Dequantize(HiPrim): def __init__(self, q_aval): self.in_avals = (q_aval,) self.out_aval = ShapedArray(q_aval.shape, jnp.dtype('float32'), sharding=q_aval.sharding) self.params = {} super().__init__() def expand(self, qx): return qx.qvalue.astype('float32') * qx.scale[..., None] def vjp_fwd(self, nzs_in, qx): return self(qx), None def vjp_bwd_retval(self, _res, g): return (g,) def quantize(x): return Quantize(jax.typeof(x))(x) def dequantize(qx): return Dequantize(jax.typeof(qx))(qx)Quantize的out_aval和Dequantize的in_avals是QArrayTy:新类型出现在原语类型签名里,和数组类型待遇相同。expand可以自由构造和检查QArray值类,因为原语实现处于抽象边界之内。
eager 执行验证
文档首先验证一切在 eager 模式下可用:
x = jnp.array([[1., 2., 3.], [4., -5., 6.]]) qx = quantize(x) print(qx) print(jax.typeof(qx)) print(dequantize(qx))成功条件:quantize(x)返回QArray,jax.typeof(qx)返回QArrayTy(借助str_short打印为q8[2,3]),dequantize(qx)返回f32数组。再确认 hi 类型确实进入了 jaxpr:
jax.jit(lambda x: dequantize(quantize(x))).trace(x).jaxpr文档说明:tracing 时量化数组显示为单一值、单一类型q8[2,3],由一条方程生产、另一条方程消费;hi 类型只在 lowering 阶段消失,那时expand被 trace,每个q8[...]类型的值按lo_ty展开成数组组件。相比之下 pytree 方案会把同一计算显示成四个看不出配对关系的数组中间值。
自定义切线类型:让梯度流过量化
切线类型是 pytree 给不了的核心能力。类型上的to_tangent_aval声明“量化数组的切线是普通f32数组”,再配合原语上的 straight-through VJP 规则,梯度就像量化是恒等函数一样流过去:
def f(x): return jnp.sum(dequantize(quantize(x))) print(jax.grad(f)(x))对量化数组输入求导时,切线类型的效果直接体现在结果类型上——梯度是普通浮点数组:
def g(qx): return jnp.sum(dequantize(qx) ** 2) print(jax.grad(g)(qx)) print(jax.typeof(jax.grad(g)(qx)))文档同时指出:把切线类型选成f32数组是一个选择。你也可以让QArrayTy的切线类型就是QArrayTy本身(切线和余切都被量化,适合不同的应用场景);做了这个选择后,由于切线类型本身是 hi 类型,还需要在该类型上实现vspace_zero和vspace_add,让 autodiff 能实例化和累加余切。
自定义批量规则:MappingSpec、dec_rank与batch
对数组,vmap的in_axes/out_axes是轴索引,JAX 能从参数形状推断被映射的轴大小。对一般 hi 类型,JAX 不做猜测:你定义一个 “mapping spec” 类型来说明你的类型如何被映射,用户把它作为in_axes/out_axes条目传入,并且当轴大小无法从数组参数推断时,显式传入axis_size。
对按行量化的QArray,一批QArray就是更大的QArray:把n个q8[2,3]沿新前导轴堆叠得到q8[n,2,3](qvalue形状(n,2,3),scale形状(n,2))。所以唯一需要的映射概念是“前导轴”,spec 类型不用携带任何数据:
from jax.experimental.hijax import MappingSpec @dataclass(frozen=True) class QArraySpec(MappingSpec): pass # QArrays are only mapped along their leading axis类型上实现dec_rank和inc_rank——hi 类型版的“去掉被映射轴”和“加上被映射轴”。它们接收轴大小和 spec,分别返回元素类型和批量化后的类型:
def qarray_dec_rank(self, size, spec): assert isinstance(spec, QArraySpec) and self.shape[0] == size return QArrayTy(self.shape[1:], self.sharding.update(spec=jax.P(*self.sharding.spec[1:]))) def qarray_inc_rank(self, size, spec): assert isinstance(spec, QArraySpec) return QArrayTy((size, *self.shape), self.sharding.update(spec=jax.P(None, *self.sharding.spec))) QArrayTy.dec_rank = qarray_dec_rank QArrayTy.inc_rank = qarray_inc_rank(文档注释:这里按 notebook 风格给类补方法;在正式代码中,它们应该直接写进class QArrayTy的定义里。)
原语上实现batch规则。规则收到批量化后的参数和它们的映射 spec(未批量的参数是None,批量化的数组参数是整数轴,批量化的 hi 类型参数是 spec 实例),返回批量化结果及其 spec。文档强调:规则必须准备好处理任意“批量/未批量”参数组合:
def quantize_batch(self, axis_data, args, in_dims): x, = args d, = in_dims if d is None: return quantize(x), None x = jnp.moveaxis(x, d, 0) return quantize(x), QArraySpec() Quantize.batch = quantize_batch def dequantize_batch(self, axis_data, args, in_dims): qx, = args d, = in_dims if d is None: return dequantize(qx), None assert isinstance(d, QArraySpec) return dequantize(qx), 0 Dequantize.batch = dequantize_batch因为按行量化在任何 rank 都成立,两条规则都可以把未批量的操作直接应用到堆叠后的值上——文档称之为“批量是同类型族成员的类型”共有的简化。
vmap 验证
映射到量化数组输出:轴大小照旧从数组参数推断,out_axes传 spec:
xs = jnp.arange(24., dtype='float32').reshape(4, 2, 3) qxs = jax.vmap(quantize, out_axes=QArraySpec())(xs) print(jax.typeof(qxs)) print(qxs.qvalue.shape, qxs.scale.shape)映射过量化数组输入:in_axes传 spec,且因为没有可推断轴大小的数组参数,必须显式传axis_size:
xs_roundtrip = jax.vmap(dequantize, in_axes=QArraySpec(), axis_size=4)(qxs) print(jax.typeof(xs_roundtrip))常规组合同样工作——vmapofjit:
print(jax.typeof(jax.vmap(jax.jit(dequantize), in_axes=QArraySpec(), axis_size=4)(qxs)))以及vmapofgrad:
def norm_quantized(x): return jnp.sum(dequantize(quantize(x)) ** 2) print(jax.vmap(jax.grad(norm_quantized))(xs).shape)容易踩的坑:容器操作只允许在抽象边界内
文档专门用两个反例划出边界:QArray的直接构造和属性读取只允许发生在expand(以及类型自己的方法,如lower_val、raise_val)里。在任何可能被jit、微分或vmap的函数中,hi 值必须只通过原语生产和消费。原因是在 trace 之下,量化数组不再是QArray实例,而是类型为q8[...]的Tracer。
在 trace 代码里读属性会直接失败:
try: jax.jit(lambda qx: qx.qvalue)(qx) except AttributeError as e: print('AttributeError:', e)更隐蔽的错误是在 trace 后的数组上调用构造函数:它不会立刻报错,而是把Tracer偷运进一个 JAX 视为不透明具体值的容器,错误在远离原因的地方才暴露——这里是 missing constant handler 的TypeError(在grad下则是 leaked-tracer 错误):
def bad_quantize(x): scale = jnp.max(jnp.abs(x), axis=-1) / 127. return QArray(jnp.round(x / scale[..., None]).astype('int8'), scale) try: jax.jit(bad_quantize)(x) except TypeError as e: print('TypeError:', e)而expand内部之所以可以直接操纵容器:等expand运行时,JAX 已经确定按该类型的 lojax 组件实现原语,它的QArray参数是真正的QArray实例(持有 lojax 值,可能是 traced 的)。文档提醒:推导规则(VJP/batch 等)是普通 traced 代码,需要访问组件时应走原语(如dequantize)而不是读属性。顶层对具体值直接看属性(如qxs.qvalue.shape)没问题,因为那是 eager 执行。
边界与下一步
- 实验状态:hijax 的导入来自
jax.experimental.hijax,文档明确 API 会演进;jax.custom_jvp/jax.custom_vjp仍是完全支持的经典工具,简单场景下可能更方便。 scan支持:jax.lax.scan始终沿前导轴遍历,因此类型上只需额外实现leading_axis_spec(返回前导轴对应的 mapping spec),dec_rank/inc_rank完成其余工作。当被扫描的值全是 hi 类型时没有可推断的长度,要显式传length。- 显式 sharding 模式:
QArrayTy的sharding字段和原语 typing 规则中的传播逻辑,是为了让 hi 类型参与显式 sharding 模式(sharding 是数组类型的一部分,jax.typeof报告跨 mesh 的划分)。没有 mesh 时这些 sharding 都是平凡的,多余代码不起作用。跨shard_map边界还需要一个HiPspec子类和类型上的shard/unshard方法。注意文档给出的一个限制:JAX 不会拿 hi 原语声明的输出 sharding 与其expand实际产生的东西交叉核对,保持二者一致是 typing 规则自己的责任。 - 更多示例:docs/301/hijax-types.md 还给出 rank-1 矩阵和通用 tuple 两个完整示例(后者展示了 spec 可以携带每组件一个轴项这类更丰富的设计);tests/hijax_test.py 被文档指为更多实例的来源;原语侧 API(JVP 规则、符号零、自定义线性化等)的深入内容见 docs/301/custom-derivatives.md。
【免费下载链接】jaxComposable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考