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

自定义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规则需要:

  1. 识别输入的batch维度(bd_x)
  2. 对每个batch切片内的元素执行求和操作(即对除batch维度外的所有维度求和)
  3. 返回保留batch维度的结果,并告知JAX结果的batch维度位置

这样jax.vmap才能正确地将自定义Primitive的逻辑映射到每个batch元素上,和原生jnp.sum的vmap行为保持一致。

内容的提问来源于stack exchange,提问作者Yahya

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.12 05:12:45