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

为不同形状尺寸数组定义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)

方案说明

  1. 单元素逻辑锚定:先实现单个数值的计算逻辑,确保jax.lax.switch能正确匹配对应函数,避免维度混乱;
  2. 分层vmap映射:
    • 第一层针对u的一维数组,将计算扩展到所有u元素;
    • 第二层针对y和z的第二个维度,覆盖所有k位置的计算;
    • 第三层针对i、x以及y/z的第一个维度,覆盖所有j位置的计算;
  3. 轴对齐控制:通过in_axes参数明确每个输入在vmap中是否被映射(0表示映射对应轴,None表示保持维度不变),确保各维度严格对齐。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.31 23:15:33