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

基于JAX优化大尺寸二维数组的行列约束选值问题

高效JAX实现行列约束的Top-K选值操作

问题回顾

给定N×N二维数组,需为每行选取数值最大的4个非零元素,在同维度矩阵C的对应位置设1,约束条件为:

  • 每行最多含4个1
  • 每列最多含4个1,列满4个1后后续行不可再选该列元素

现有Python循环+字典的实现处理4000×4000矩阵耗时超10分钟,核心问题是Python循环无法利用JAX的向量化和JIT编译优化,且频繁的数组修改带来大量内存开销。

优化思路

  1. 替换Python循环为JAX扫描操作:用jax.lax.scan实现状态迭代,将循环逻辑编译为设备端高效代码,避免Python解释器开销。
  2. 用数组替代字典跟踪列计数:JAX数组支持向量化操作,比字典的逐元素查找修改效率高得多,且可被JIT编译。
  3. 预处理排序减少重复计算:一次性完成所有行的降序索引排序,避免循环中重复排序的冗余操作。
  4. 批量更新状态:通过JAX的at索引API批量更新列计数和C矩阵,减少数组复制次数。

实现代码

import jax
import jax.numpy as jnp

@jax.jit
def constrained_top_k(a, k=4):
    n = a.shape[0]
    # 预处理:每行按元素值降序排列的列索引
    sorted_row_indices = jnp.flip(jnp.argsort(a, axis=1), axis=1)
    # 初始化状态:列计数数组、结果矩阵C
    init_state = (
        jnp.zeros(n, dtype=jnp.int32),  # 记录每列已选元素数量
        jnp.zeros((n, n), dtype=jnp.int32)  # 结果矩阵C
    )

    def scan_step(state, inputs):
        row_idx, row_sorted_cols = inputs
        col_counts, c_matrix = state

        # 筛选当前行的有效候选列:列计数未达上限、对应元素非零
        valid_mask = (col_counts[row_sorted_cols] < k) & (a[row_idx, row_sorted_cols] != 0)
        valid_cols = row_sorted_cols[valid_mask]
        # 取最多k个有效列
        selected_cols = valid_cols[:k]

        # 更新状态:列计数+1、结果矩阵设1
        new_col_counts = col_counts.at[selected_cols].add(1)
        new_c_matrix = c_matrix.at[row_idx, selected_cols].set(1)
        return (new_col_counts, new_c_matrix), None

    # 遍历所有行执行扫描
    scan_inputs = (jnp.arange(n), sorted_row_indices)
    final_state, _ = jax.lax.scan(scan_step, init_state, scan_inputs)
    return final_state[1]

# 测试示例
a_array = jnp.array([[0,3,4,0,0,12,19,22],
            [7,0,0,10,0,0,0,15],
            [12,0,0,15,16,19,0,31],
            [17,18,0,0,21,23,78,89],
            [22,2,78,0,0,1111,12,33],
            [123,0,122,10,14,0,50,60],
            [10,110,0,1231,0,110,0,61],
            [0,17,0,141,0,166,16,0]])

c_matrix = constrained_top_k(a_array)
print(c_matrix)

性能说明

  • 对于4000×4000矩阵,JIT编译后的constrained_top_k函数可将耗时压缩至秒级(具体取决于硬件),远优于原Python循环实现。
  • 若矩阵中0元素占比极高,可提前生成非零元素的掩码数组,进一步减少筛选时的计算量。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.20 20:35:11