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

如何在JAX中高效按索引将矩阵元素累加至目标矩阵

JAX 无循环实现分箱求和

针对你的需求,JAX提供了两种高效的无循环实现方式,可完全替代逐元素循环的累加操作,且能轻松扩展到更高维度:

方法1:向量化使用 jnp.add.at

jnp.add.at 本身支持批量索引操作,无需手动遍历每个元素。你可以直接传入所有索引和对应的值,它会自动完成分组累加:

import jax.numpy as jnp
import numpy as np

values = np.random.random((1000,2)) * 100
bins = jnp.linspace(0,100,1000)
indices = jnp.digitize(values, bins)
sol = jnp.zeros((len(bins), len(bins), 2))

# 向量化更新,替代原循环
sol = sol.at[indices[:,0], indices[:,1]].add(values)

原理说明

  • indices[:,0] 和 indices[:,1] 分别提取所有元素的行、列索引,形成两个1维数组
  • jnp.add.at 会自动将相同 (行索引, 列索引) 位置的 values 元素累加,逻辑和原循环完全一致,但底层是向量化操作,性能远高于Python循环,且天然支持高维扩展

方法2:使用 jax.ops.segment_sum 实现分组求和

如果需要更灵活的分组逻辑(尤其是高维场景),可以用 segment_sum。需先将多维索引转换为一维扁平索引,再进行分组求和:

import jax.numpy as jnp
import numpy as np

values = np.random.random((1000,2)) * 100
bins = jnp.linspace(0,100,1000)
indices = jnp.digitize(values, bins)
sol_shape = (len(bins), len(bins), 2)

# 将二维索引转换为一维扁平索引
flat_indices = indices[:,0] * sol_shape[1] + indices[:,1]

# 按扁平索引分组求和,再重塑为目标形状
sol = jax.ops.segment_sum(values, flat_indices, num_segments=sol_shape[0]*sol_shape[1])
sol = sol.reshape(sol_shape)

原理说明

  • 把二维的 (i,j) 索引转换成一维的 i*cols + j,让每个二维位置对应唯一的一维分组ID
  • segment_sum 会对每个分组ID对应的 values 元素求和,最后将结果重塑回目标的多维形状
  • 这种方式在高维场景下扩展更方便,比如三维索引可转换为 i*cols*depth + j*depth + k,逻辑统一

性能对比

两种方法均远优于Python循环:

  • jnp.add.at 语法简洁,和原循环逻辑最贴近,适合快速替换
  • segment_sum 在数据量极大时性能更优,且分组逻辑更通用

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.18 05:15:02