如何用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
相关产品推荐
相关产品推荐

