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

Jax中vmapped函数索引报错咨询:jnp.take无法提取坐标

解决Jax中vmap+索引操作的IndexError问题

错误根源

  1. 混用vectorize和vmap:你给main加了@jnp.vectorize装饰器,又用双重vmap调用compute_S,导致输入的P、M被过度拆解,进入main时变成标量而非(3,)的点向量,所以jnp.take会报索引过多的错误。
  2. 索引方式冗余:提取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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.23 07:35:34