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

如何用JAX高效计算矩阵下三角元素与向量元素对运算

如何高效计算向量所有j<i元素对的指定值并输出下三角展平结果

请问是否可以高效对向量的所有元素对执行指定计算?我的需求是填充矩阵的下三角元素(支持展平为一维结构)。

具体需求

  • 对所有索引范围在[1, length(input_vector)]区间内、且满足j < i的索引对(i,j),计算do_my_calculation(input_vector[i], input_vector[j])
  • 保存全部计算结果

我对结果的形状没有严格要求,如果可以选择的话,优先返回下三角(i,j)矩阵元素按顺序展平得到的一维向量。

目标逻辑伪代码

input_vector = np.arange(100)
result_vector = []

for i in range(1, len(input_vector)):
    for j in range(0, i):
        result_vector.append(do_my_calculation(input_vector[i], input_vector[j]))

注:上述代码中input_vector和result_vector的类型不做限定,如有需要完全可以预分配result_vector的存储空间,示例中使用列表仅为简化示例代码。


补充:具体实现示例

注:本问题的核心并非如何在JAX中跑通该逻辑,而是如何实现高效向量化运行,避免慢速的显式Python循环。

# 问题初始化
import numpy as np

dim = 15
input_vector_x = np.random.rand(dim)
input_vector_y = np.random.rand(dim)
output_vector = np.empty(np.tril_indices(dim, k=-1)[0].size)

assert input_vector_x.size == input_vector_y.size

# 实现方式1:显式双层循环(运行效率低)
counter = 0
for i in range(1, input_vector_x.size):
    for j in range(0, i):
        output_vector[counter] = (input_vector_y[j] - input_vector_y[i]) / (input_vector_x[j] - input_vector_x[i])
        counter += 1

# 实现方式2:索引向量化实现(运行效率高)
indices = np.tril_indices(dim, k=-1)
i = indices[0]
j = indices[1]

output_vector = (input_vector_y[j] - input_vector_y[i]) / (input_vector_x[j] - input_vector_x[i])

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.01 19:54:20