如何向量化包含多索引访问的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
相关产品推荐
相关产品推荐

