为何无法用vmap向量化jax.lax.cond处理布尔数组?
问题原因与解决方案
你的代码无法得到预期数组的核心问题是jax.lax.cond的函数签名不匹配,以及对其参数传递逻辑的误解。
问题分析
jax.lax.cond的调用规则是:
- 当传入
pred(布尔标量)后,true_fun和false_fun的参数数量必须与operand匹配:- 如果省略
operand(或传None),两个函数必须是无参函数; - 如果指定
operand,两个函数必须接收一个参数(即该operand的值)。
- 如果省略
你之前的代码中,f1 = lambda x:1和f2 = lambda y:0定义了带参数的函数,但调用jax.lax.cond(dk,f1,f2)时没有传入operand,这会导致JAX尝试给无参的调用场景传递参数,引发签名不匹配的错误,自然无法得到预期的数组输出。
修正后的代码
方案1:使用无参函数(适合不需要用到输入值的场景)
import jax.numpy as jnp import jax DK = jnp.array([[True],[True],[False],[True]]) f1 = lambda: 1 # 定义无参函数 f2 = lambda: 0 cond = lambda dk: jax.lax.cond(dk, f1, f2) vcond = jax.vmap(cond) print(vcond(DK)) # 输出:[[1] # [1] # [0] # [1]]
方案2:显式传递operand(适合后续需要基于输入值计算的场景)
如果后续需要在f1/f2中使用输入的布尔值或其他数据,可以显式指定operand:
import jax.numpy as jnp import jax DK = jnp.array([[True],[True],[False],[True]]) # 函数接收operand参数,这里忽略参数直接返回固定值,也可以用参数做计算 f1 = lambda x: 1 f2 = lambda x: 0 # 显式传递operand,这里传dk本身(也可以传其他值) cond = lambda dk: jax.lax.cond(dk, f1, f2, operand=dk) vcond = jax.vmap(cond) print(vcond(DK)) # 输出与方案1一致
额外优化:输出一维数组
如果希望输出一维数组而非二维,可以对输入做squeeze处理,或者指定vmap的映射维度:
# 方式1:压缩输入的维度 print(vcond(DK.squeeze())) # 输出:[1 1 0 1] # 方式2:指定vmap的in_axes映射第二个维度 vcond_ax1 = jax.vmap(cond, in_axes=1) print(vcond_ax1(DK).squeeze()) # 输出:[1 1 0 1]
内容的提问来源于stack exchange,提问作者Kapil
相关产品推荐
相关产品推荐

