自定义PyTree的aux_data为何在jax.jit后对jnp.array被追踪而非np.array?
PyTree辅助数据在JAX JIT下的追踪行为解析
我正在研究PyTree的工作机制,将自定义类注册为PyTree时发现一个特性:
- 当PyTree的aux_data(静态数据)是
jax.numpy.ndarray时,经过jax.jit()转换后,该辅助数据会被追踪,返回Traced<ShapedArray(...)>类型; - 若aux_data是普通
numpy.ndarray,则不会被追踪,JIT转换后的函数直接返回原数组。
我了解JAX JIT的追踪机制,但从PyTree层面无法理解这种差异。我用equinox、simple_pytree等成熟库的自定义PyTree实现做了对比测试,结果完全一致,说明这是JAX PyTree的设计特性而非Bug。
复现代码
import jax from jax.tree_util import tree_structure, tree_leaves import numpy as np def get_pytree_impl(base): if base == "equinox": import equinox as eqx Module = eqx.Module static_field = eqx.static_field elif base == "simple_pytree": from simple_pytree import Pytree, static_field Module = Pytree elif base == "dataclasses": from dataclasses import dataclass, field @dataclass class Module(): pass static_field = field class PytreeImpl(Module): x: jax.numpy.ndarray y: jax.numpy.ndarray = static_field() def __init__(self, x, y): self.x = x self.y = y if base == 'dataclasses': from jax.tree_util import register_pytree_node def flatten(ptree): return ((ptree.x,), ptree.y) def unflatten(aux_data, children): return PytreeImpl(*children, aux_data) register_pytree_node(PytreeImpl, flatten, unflatten) return PytreeImpl def times_two(ptree): return type(ptree)(ptree.x*2, ptree.y*2) times_two_jitted = jax.jit(times_two) bases = ['dataclasses', 'equinox', 'simple_pytree'] for base in bases: print("======== " + base + " ========") for lib_name, array_lib in zip(['jnp', 'np'], [jax.numpy, np]): print("==== " + lib_name) PytreeImpl = get_pytree_impl(base) x = jax.numpy.array([1,2]) y = array_lib.array([3,4]) input_tree = PytreeImpl(x, y) for tag, pytree in zip(["input", "no_jit", "jit"],[input_tree, times_two(input_tree), times_two_jitted(input_tree)]): print(f' {tag}:') print(f'\t Structure: {tree_structure(pytree)}') print(f'\t Leaves: {tree_leaves(pytree)}')
运行输出
======== dataclasses ======== ==== jnp input: Structure: PyTreeDef(CustomNode(PytreeImpl[[3 4]], [*])) Leaves: [Array([1, 2], dtype=int32)] no_jit: Structure: PyTreeDef(CustomNode(PytreeImpl[[6 8]], [*])) Leaves: [Array([2, 4], dtype=int32)] jit: Structure: PyTreeDef(CustomNode(PytreeImpl[Traced<ShapedArray(int32[2])>with<DynamicJaxprTrace(level=1/0)>], [*])) Leaves: [Array([2, 4], dtype=int32)] ==== np input: Structure: PyTreeDef(CustomNode(PytreeImpl[[3 4]], [*])) Leaves: [Array([1, 2], dtype=int32)] no_jit: Structure: PyTreeDef(CustomNode(PytreeImpl[[6 8]], [*])) Leaves: [Array([2, 4], dtype=int32)] jit: Structure: PyTreeDef(CustomNode(PytreeImpl[[6 8]], [*])) Leaves: [Array([2, 4], dtype=int32)] ======== equinox ======== ==== jnp input: Structure: PyTreeDef(CustomNode(PytreeImpl[('x',), ('y',), (Array([3, 4], dtype=int32),)], [*])) Leaves: [Array([1, 2], dtype=int32)] no_jit: Structure: PyTreeDef(CustomNode(PytreeImpl[('x',), ('y',), (Array([6, 8], dtype=int32),)], [*])) Leaves: [Array([2, 4], dtype=int32)] jit: Structure: PyTreeDef(CustomNode(PytreeImpl[('x',), ('y',), (Traced<ShapedArray(int32[2])>with<DynamicJaxprTrace(level=1/0)>),)], [*])) Leaves: [Array([2, 4], dtype=int32)] ==== np input: Structure: PyTreeDef(CustomNode(PytreeImpl[('x',), ('y',), (array([3, 4]),)], [*])) Leaves: [Array([1, 2], dtype=int32)] no_jit: Structure: PyTreeDef(CustomNode(PytreeImpl[('x',), ('y',), (array([6, 8]),)], [*])) Leaves: [Array([2, 4], dtype=int32)] jit: Structure: PyTreeDef(CustomNode(PytreeImpl[('x',), ('y',), (array([6, 8]),)], [*])) Leaves: [Array([2, 4], dtype=int32)] ======== simple_pytree ======== ==== jnp input: Structure: PyTreeDef(CustomNode(PytreeImpl[(('x',), {'y': Array([3, 4], dtype=int32), '_pytree__initialized': True})], [*])) Leaves: [Array([1, 2], dtype=int32)] no_jit: Structure: PyTreeDef(CustomNode(PytreeImpl[(('x',), {'y': Array([6, 8], dtype=int32), '_pytree__initialized': True})], [*])) Leaves: [Array([2, 4], dtype=int32)] jit: Structure: PyTreeDef(CustomNode(PytreeImpl[(('x',), {'y': Traced<ShapedArray(int32[2])>with<DynamicJaxprTrace(level=1/0)>, '_pytree__initialized': True})], [*])) Leaves: [Array([2, 4], dtype=int32)] ==== np input: Structure: PyTreeDef(CustomNode(PytreeImpl[(('x',), {'y': array([3, 4]), '_pytree__initialized': True})], [*])) Leaves: [Array([1, 2], dtype=int32)] no_jit: Structure: PyTreeDef(CustomNode(PytreeImpl[(('x',), {'y': array([6, 8]), '_pytree__initialized': True})], [*])) Leaves: [Array([2, 4], dtype=int32)] jit: Structure: PyTreeDef(CustomNode(PytreeImpl[(('x',), {'y': array([6, 8]), '_pytree__initialized': True})], [*])) Leaves: [Array([2, 4], dtype=int32)]
行为解析
PyTree的aux_data属于静态结构信息,JIT编译时会将其视为不会随输入变化的固定数据,但针对不同类型的数组,JAX的处理逻辑有差异:
- jax.numpy.ndarray作为aux_data:JAX会对静态数据中的JAX数组进行形状追踪——即使是静态数据,JAX也需要明确其形状来生成合法的jaxpr(JAX的中间表示),因此会用
Traced<ShapedArray>对象包裹,只保留形状信息而不追踪具体数值。 - numpy.ndarray作为aux_data:普通numpy数组不属于JAX的原生可追踪类型,JIT会直接将其当作Python原生对象处理,不会进行任何追踪,因此返回时保持原数组的数值与类型。
另外需要注意:示例中在JIT函数里修改aux_data(y*2)的操作其实不符合静态数据的设计初衷——静态数据本应在JIT编译阶段就被固定,运行时不会改变。这里的修改能生效是因为JIT对静态数据的处理逻辑允许在追踪阶段执行一次计算,但这种写法不推荐,可能导致意外行为。
环境依赖:
- Python 3.12.1
- equinox 0.11.4
- jax 0.4.28
- jaxlib 0.4.28
- simple-pytree 0.1.5
内容的提问来源于stack exchange,提问作者alx
相关产品推荐
相关产品推荐

