如何为多输出场景编写正确的JAX Batching Rule?
JAX自定义Primitive多输出批处理规则修正方案
问题核心
你的自定义Primitive批处理规则存在两个关键问题:
- 第二个输出未被正确批处理:错误标记了其batch维度为输入的batch轴,且未让JAX自动处理常量输出的广播逻辑;
out_axes配置报错:由于batch维度标记错误,导致JAX无法匹配你设置的输出轴规则。
修正后的完整代码
import jax import jax.numpy as jnp from jax.extend import core from jax.interpreters import batching my_sum_p = core.Primitive('my_sum') my_sum_p.multiple_results = True def my_sum(x): return my_sum_p.bind(x) # 原语求值规则保持不变 def my_sum_impl(x): return jnp.sum(x), jnp.ones((3,)) my_sum_p.def_impl(my_sum_impl) # 修正后的批处理规则 def my_sum_batching_rule(batched_args, batch_dims): x, = batched_args bd_x, = batch_dims # 处理第一个输出:对输入除batch轴外的所有轴求和,保留batch轴 axis = tuple(i for i in range(x.ndim) if i != bd_x) sum_output = jnp.sum(x, axis=axis) # 处理第二个输出:返回原常量,不依赖输入batch维度 ones_output = jnp.ones((3,)) # 标记输出的batch维度:第一个输出对应输入的batch轴,第二个输出无batch轴 return (sum_output, ones_output), (bd_x, None) batching.primitive_batchers[my_sum_p] = my_sum_batching_rule # 测试验证 if __name__ == "__main__": x = jnp.array([[1.0, 2.0, 3.0], [4.0, 5.0, 6.0]]) # 默认vmap配置,输出与原函数一致 vms_default = jax.vmap(my_sum) print('my_sum 默认vmap:', vms_default(x)) # 设置out_axes=(0, None),正常运行 vms_custom = jax.vmap(my_sum, out_axes=(0, None)) print('my_sum 自定义out_axes:', vms_custom(x)) # 对比原函数输出 def original_sum(x): return jnp.sum(x), jnp.ones((3,)) vos_default = jax.vmap(original_sum) print('original_sum 默认vmap:', vos_default(x)) vos_custom = jax.vmap(original_sum, out_axes=(0, None)) print('original_sum 自定义out_axes:', vos_custom(x))
关键修正说明
- 输出batch维度标记:将第二个输出的batch维度标记为
None,明确告诉JAX该输出是不依赖输入batch的常量,JAX会自动根据out_axes配置处理广播或形状保留; - 输出生成逻辑:第二个输出直接返回原常量即可,无需手动扩展batch维度——JAX会根据配置自动处理:
- 默认
out_axes=(0,0)时,自动广播到batch规模,得到(batch_size, 3)的数组; - 设置
out_axes=(0, None)时,保留原输出形状(3,),不会触发维度不匹配错误。
- 默认
运行结果
my_sum 默认vmap: (Array([ 6., 15.], dtype=float32), Array([[1., 1., 1.], [1., 1., 1.]], dtype=float32)) my_sum 自定义out_axes: (Array([ 6., 15.], dtype=float32), Array([1., 1., 1.], dtype=float32)) original_sum 默认vmap: (Array([ 6., 15.], dtype=float32), Array([[1., 1., 1.], [1., 1., 1.]], dtype=float32)) original_sum 自定义out_axes: (Array([ 6., 15.], dtype=float32), Array([1., 1., 1.], dtype=float32))
内容的提问来源于stack exchange,提问作者Yahya
相关产品推荐
相关产品推荐

