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

如何向量化包含多索引访问的Numpy for循环以处理unigram数组计算

大数组循环计算向量化及加速需求
  • unigram为形状(N, M, 100)的数组
  • 需完全去掉for循环完成全部计算
  • seq是长度为M的一维数组,M最高可达10000

原有循环实现代码

batch_size, seq_len, num_labels = unigram_scores.shape
broadcast = np.broadcast_to(seq, (batch_size, seq_len))


for i in range(0, broadcast.shape[1]):
    n_seq[i] = unigram_scores[np.arange(batch_size), i , broadcast[:,i]]

已尝试的优化方案

  • 原生NumPy实现:@hpaulj 给出的方案可正常运行,无需额外安装依赖,但实际提速效果远低于预期
  • Numba并行优化:实现代码如下
import numpy as np
from numba import njit, prange

@njit(parallel=True)
def calculate_unigram_probability(unigram_scores,seq):
    
    batch_size, seq_len, num_labels = unigram_scores.shape
    broadcast = np.broadcast_to(seq, (batch_size, seq_len))

    for i in prange( broadcast.shape[1]):
       n_seq[i] = unigram_scores[np.arange(batch_size), i , broadcast[:,i]]


    return n_seq

上述Numba优化方案运行速度仍然偏慢,目前正在尝试将计算从CPU迁移到CUDA,以达到预期的提速效果。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.30 22:18:04