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

如何为多输出场景编写正确的JAX Batching Rule?

JAX自定义Primitive多输出批处理规则修正方案

问题核心

你的自定义Primitive批处理规则存在两个关键问题:

  1. 第二个输出未被正确批处理:错误标记了其batch维度为输入的batch轴,且未让JAX自动处理常量输出的广播逻辑;
  2. 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))

关键修正说明

  1. 输出batch维度标记:将第二个输出的batch维度标记为None,明确告诉JAX该输出是不依赖输入batch的常量,JAX会自动根据out_axes配置处理广播或形状保留;
  2. 输出生成逻辑:第二个输出直接返回原常量即可,无需手动扩展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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.12 05:14:50