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

能否让jax.vmap()直接实现hstack()操作以规避拷贝瓶颈?

问题1:能否让jax.vmap直接完成hstack操作避免拷贝?

不行,jax.vmap的设计是保持单步函数的输出结构,批量映射输入轴,无法直接让它输出hstack后的结果。但可以通过调整数组维度变换的方式,替代显式的hstack操作,利用JAX的融合优化减少额外开销:

原vmap输出是形状为(m, 2, 2)的数组,我们可以通过转置+reshape直接得到目标形状:

desired_output = vmap_output.transpose(1, 0, 2).reshape(2, -1)

这个操作会被JAX的XLA编译器融合成单一操作,避免hstack带来的额外拷贝开销,效果和jnp.hstack(vmap_output)完全一致,但性能更优。

验证代码:

import jax
import jax.numpy as jnp

def f(a, b, c):
  return jnp.array([[a.sum(), b.sum()], [c.sum(), 0.]]) # 返回2x2数组

def arr(m, n):
  return jnp.arange(m*n).reshape((m, n))

m = 3

a = arr(m, 2)
b = arr(m, 5)
c = arr(m, 7)

fv = jax.vmap(f)
vmap_output = fv(a, b, c)

# 替代hstack的高效方式
desired_output = vmap_output.transpose(1, 0, 2).reshape(2, -1)

print(jnp.allclose(desired_output, jnp.hstack(vmap_output))) # 输出True
问题2:动态更新切片未保留结果的解决办法

你的代码问题在于JAX是函数式编程范式,所有数组都是不可变的:jax.lax.dynamic_update_slice_in_dim不会原地修改输入数组,而是返回一个新的修改后数组,但你的函数g没有返回这个结果;同时vmap并行处理每个索引时,无法自动累积更新(每个vmap分支都是独立的)。

方法1:用scan累积更新预分配数组

def g(carry, idx):
  a_slice = a[idx]
  b_slice = b[idx]
  c_slice = c[idx]
  block = jnp.array([[a_slice.sum(), b_slice.sum()], [c_slice.sum(), 0.]])
  updated_carry = jax.lax.dynamic_update_slice_in_dim(carry, block, idx*2, axis=1)
  return updated_carry, None

g_output = jnp.zeros((2, 2*m))
final_output, _ = jax.lax.scan(g, g_output, jnp.arange(m))

print(final_output)
# 输出:
# [[  1.  10.   5.  35.   9.  60.]
#  [ 21.   0.  70.   0. 119.   0.]]

方法2:直接生成所有block后拼接(更高效)

回到最初的思路,用vmap生成所有block后,通过转置+reshape得到结果,比手动更新更简洁高效,JAX会自动优化整个流程:

vmap_output = fv(a, b, c)
final_output = vmap_output.transpose(1, 0, 2).reshape(2, -1)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.29 11:53:12