自定义JAX Primitive的Batching Rule编写错误,结果不符
修复JAX自定义Primitive的vmap行为不一致问题
你的自定义my_sum Primitive的batching规则逻辑错误,导致jax.vmap调用时没有按预期对每个batch元素单独求和,而是直接计算了全局总和。
修正后的代码
import jax import jax.numpy as jnp from jax.extend import core from jax.interpreters import batching my_sum_p = core.Primitive('my_sum') def my_sum(x): return my_sum_p.bind(x) # 原语求值规则 def my_sum_impl(x): return jnp.sum(x) my_sum_p.def_impl(my_sum_impl) # 正确的batching规则 def my_sum_batching_rule(batched_args, batch_dims): x, = batched_args bd_x, = batch_dims # 对输入中除batch维度外的所有维度求和,保留batch维度 non_batch_axes = tuple(d for d in range(x.ndim) if d != bd_x) result = jnp.sum(x, axis=non_batch_axes) # 返回结果及结果的batch维度位置 return result, bd_x 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]]) vms = jax.vmap(my_sum) print('my_sum:', vms(x)) # 输出: [ 6. 15.] def original_sum(x): return jnp.sum(x) vos = jax.vmap(original_sum) print('original_sum:', vos(x)) # 输出: [ 6. 15.]
问题原因解释
原batching规则中直接调用my_sum(x),相当于对整个输入数组求和,完全忽略了batch维度的存在。而正确的batching规则需要:
- 识别输入的batch维度(
bd_x) - 对每个batch切片内的元素执行求和操作(即对除batch维度外的所有维度求和)
- 返回保留batch维度的结果,并告知JAX结果的batch维度位置
这样jax.vmap才能正确地将自定义Primitive的逻辑映射到每个batch元素上,和原生jnp.sum的vmap行为保持一致。
内容的提问来源于stack exchange,提问作者Yahya
相关产品推荐
相关产品推荐

