为不同形状尺寸数组定义JAX vmap的正确向量化轴
基于JAX动态函数切换的维度适配解决方案
原有实现背景
我们基于jax.lax.switch实现了可根据索引数组动态切换多函数的f_switch函数,代码如下:
import jax import jax.random as random def g_0(x, y, z, u): return x + y + z + u def g_1(x, y, z, u): return x * y * z * u def g_2(x, y, z, u): return x - y + z - u def g_3(x, y, z, u): return x / y / z / u g_i = [g_0, g_1, g_2, g_3] @jax.jit def f_switch(i, x, y, z, u): g = lambda i: jax.lax.switch(i, g_i, x, y, z, u) return jax.vmap(g)(i)
该函数在处理shape=(len_i,)的索引数组i_ar、shape=(len_xyz,)的x_ar/y_ar/z_ar,以及shape=(len_u, len_xyz)的u_ar时,能正常输出shape=(len_i, len_xyz, len_u)的结果。
新需求场景
现在需要适配以下维度的输入数组,输出形状为(j_len, k_len, l_len)的结果:
j_len = 82 k_len = 20 l_len = 100 i_ar = random.randint(random.PRNGKey(0), shape=(j_len,), minval=0, maxval=len(g_i)) x_ar = random.uniform(random.PRNGKey(1), shape=(j_len,)) y_ar = random.uniform(random.PRNGKey(2), shape=(j_len,k_len)) z_ar = random.uniform(random.PRNGKey(3), shape=(j_len,k_len)) u_ar = random.uniform(random.PRNGKey(4), shape=(l_len,))
各输入的维度对应关系为:i_ar[j]、x_ar[j]、y_ar[j,k]、z_ar[j,k]、u_ar[l]。此前尝试嵌套vmap或广播u_ar的方案均失败,以下是正确实现方案:
正确实现方案
核心思路是通过多层vmap匹配每个维度的映射关系,明确每个输入轴的对应关系:
import jax import jax.random as random def g_0(x, y, z, u): return x + y + z + u def g_1(x, y, z, u): return x * y * z * u def g_2(x, y, z, u): return x - y + z - u def g_3(x, y, z, u): return x / y / z / u g_i = [g_0, g_1, g_2, g_3] @jax.jit def f_switch(i, x, y, z, u): # 定义单元素级别的switch计算逻辑 def single_g(i_val, x_val, y_val, z_val, u_val): return jax.lax.switch(i_val, g_i, x_val, y_val, z_val, u_val) # 第一层vmap:映射u的l_len维度 map_u = jax.vmap(single_g, in_axes=(None, None, None, None, 0)) # 第二层vmap:映射y、z的k_len维度 map_k = jax.vmap(map_u, in_axes=(None, None, 0, 0, None)) # 第三层vmap:映射i、x以及y/z的j_len维度 map_j = jax.vmap(map_k, in_axes=(0, 0, 0, 0, None)) return map_j(i, x, y, z, u) # 验证代码 j_len = 82 k_len = 20 l_len = 100 i_ar = random.randint(random.PRNGKey(0), shape=(j_len,), minval=0, maxval=len(g_i)) x_ar = random.uniform(random.PRNGKey(1), shape=(j_len,)) y_ar = random.uniform(random.PRNGKey(2), shape=(j_len,k_len)) z_ar = random.uniform(random.PRNGKey(3), shape=(j_len,k_len)) u_ar = random.uniform(random.PRNGKey(4), shape=(l_len,)) output = f_switch(i_ar, x_ar, y_ar, z_ar, u_ar) print(output.shape) # 输出 (82, 20, 100)
方案说明
- 单元素逻辑锚定:先实现单个数值的计算逻辑,确保
jax.lax.switch能正确匹配对应函数,避免维度混乱; - 分层vmap映射:
- 第一层针对
u的一维数组,将计算扩展到所有u元素; - 第二层针对
y和z的第二个维度,覆盖所有k位置的计算; - 第三层针对
i、x以及y/z的第一个维度,覆盖所有j位置的计算;
- 第一层针对
- 轴对齐控制:通过
in_axes参数明确每个输入在vmap中是否被映射(0表示映射对应轴,None表示保持维度不变),确保各维度严格对齐。
内容的提问来源于stack exchange,提问作者stefan_chem
相关产品推荐
相关产品推荐

