Jax中vmapped函数索引报错咨询:jnp.take无法提取坐标
解决Jax中vmap+索引操作的IndexError问题
错误根源
- 混用vectorize和vmap:你给
main加了@jnp.vectorize装饰器,又用双重vmap调用compute_S,导致输入的P、M被过度拆解,进入main时变成标量而非(3,)的点向量,所以jnp.take会报索引过多的错误。 - 索引方式冗余:提取x、y坐标用
jnp.take(M, [0,1])完全没必要,直接用切片更直观,还能避免维度混乱。
修正方案
1. 移除vectorize装饰器
删掉main函数上的@jnp.vectorize,让vmap单独处理批量映射逻辑,避免维度冲突。
2. 简化索引操作
把jnp.take(P, jnp.array([0, 1]))改成P[:2],同理M的部分改成M[:2],针对单个(3,)的点向量直接取前两个元素。
3. 修正后的完整代码
main函数:
def main(P, M, omega): g = 9.8 k0 = (omega**2)/g # 直接用切片提取x、y坐标 p_xy = P[:2] m_xy = M[:2] r = jnp.sqrt(jnp.sum(jnp.square(p_xy - m_xy))) Z = P[2] + M[2] # zp + zm R = jnp.sqrt(jnp.sum(jnp.square(P-M))) R1 = jnp.sqrt(jnp.square(r) + jnp.square(Z)) term = R + Z + R1 + k0 return term
调用代码:
def compute_S(P, M): omega = 0.1 val = main(P, M, omega) # 原代码里写的是main(P,P,omega),如果是笔误请修正为P和M配对 return val quad_points = jnp.array([[x1,y1,z1],[x2,y2,z2]]) # 双重vmap实现所有点对的两两计算 vectorized_pairwise_compute_S = vmap(vmap(compute_S, (0, None)), (None, 0)) S = vectorized_pairwise_compute_S(quad_points, quad_points)
额外说明
- 原调用代码里
main(P,P, omega)如果是笔误(应该是main(P,M,omega)),记得修正,否则计算的是每个点和自身的配对值。 - 双重vmap后,
vectorized_pairwise_compute_S会处理quad_points中所有点对的组合,输出形状为(n, n)的数组(n是quad_points的长度)。
内容的提问来源于stack exchange,提问作者Kapil
相关产品推荐
相关产品推荐

