You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

自定义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&lt;ShapedArray(int32[2])&gt;with&lt;DynamicJaxprTrace(level=1/0)&gt;], [*]))
     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&lt;ShapedArray(int32[2])&gt;with&lt;DynamicJaxprTrace(level=1/0)&gt;),)], [*]))
     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&lt;ShapedArray(int32[2])&gt;with&lt;DynamicJaxprTrace(level=1/0)&gt;, '_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的处理逻辑有差异:

  1. jax.numpy.ndarray作为aux_data:JAX会对静态数据中的JAX数组进行形状追踪——即使是静态数据,JAX也需要明确其形状来生成合法的jaxpr(JAX的中间表示),因此会用Traced<ShapedArray>对象包裹,只保留形状信息而不追踪具体数值。
  2. 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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.24 00:08:08