如何测试JIT编译的Jax函数是创建新张量还是视图?
判断JAX操作是创建新张量还是生成视图的方法
问题背景
我编写了如下JAX代码:
@jit def concat_permute(indices, in1, in2): tensor = jnp.concatenate([jnp.atleast_1d(in1), jnp.atleast_1d(in2)]) return tensor[indices]
测试代码:
key = jax.random.PRNGKey(758493) in1 = tens = jax.random.uniform(key, shape=(15,5,3)) in2 = tens = jax.random.uniform(key, shape=(10,5,3)) indices = jax.random.choice(key, 25, (25,), replace=False)
通过jax.make_jaxpr得到了函数的Jaxpr,但无法确定该操作是创建新张量还是生成视图,想找更清晰的判断方法。
判断方法与结论
从Jaxpr直接分析
你的Jaxpr里有两个关键操作:
concatenate:拼接操作必然会创建新的连续张量——因为两个输入数组在内存中是独立的,拼接需要将它们的数据复制到新的连续内存块中,不可能是原输入的视图。gather:按任意索引取元素的gather操作,由于索引是无序且任意排列的,无法通过原拼接张量的内存偏移直接映射结果,必须复制数据生成新张量。
所以从Jaxpr就能确定,这个函数的输出是新创建的张量,不是视图。
更直接的验证手段
1. 检查CPU数组的内存地址(仅适用于CPU)
JAX的CPU数组兼容NumPy的__array_interface__,可以通过对比内存地址判断是否为新张量:
result = concat_permute(indices, in1, in2) # 打印原输入和结果的内存起始地址 print(in1.__array_interface__['data'][0]) print(result.__array_interface__['data'][0])
如果地址不同,说明结果是新张量;视图的地址会和原数组一致(或存在固定偏移)。
2. 查看XLA的HLO编译输出
通过以下代码生成HLO文本,查看底层操作:
hlo_text = jax.jit(concat_permute).lower(indices, in1, in2).compile().as_text() print(hlo_text)
如果输出中包含copy或gather操作,都是生成新张量的明确信号——XLA的gather操作不会生成视图,必然复制数据。
3. 基于JAX不可变数组模型的常识
JAX的数组是不可变的,绝大多数重排、拼接、索引操作都会生成新张量。只有少数操作(如jnp.reshape在数据连续时、简单的切片)会生成视图,而你的操作涉及拼接+任意排列索引,完全不符合视图的生成条件。
内容的提问来源于stack exchange,提问作者nazimorhan
相关产品推荐
相关产品推荐

