如何加速大尺寸NumPy数组的循环计算?
NumPy循环计算优化方案
我正在尝试用NumPy数组执行以下循环计算,数组最大尺寸约为(5000,1000,3)。由于部分变量是动态生成的,不确定能否进一步向量化。另外,我需要用不同参数重复这个循环100多次,计算效率问题会更突出,恳请提供优化建议。
原代码
import itertools import numpy as np from numpy import random import time def get_p(payoffs): allmax = payoffs.max(axis=1)[:, None] findmax = payoffs - allmax mask = ((findmax[:, 1, :] == 0) & (findmax[:, 2, :] == 0)) findmax[:, 1, :][mask] = -1 mask = ((findmax[:, 0, :] == 0) & (findmax[:, 1, :] == 0)) findmax[:, 0, :][mask] = -1 mask = ((findmax[:, 0, :] == 0) & (findmax[:, 1, :] == 0) & (findmax[:, 2, :] == 0)) findmax[:, 0, :][mask] = -1 findmax[:, 1, :][mask] = -1 p = np.where(findmax < 0, 0.0, 1.0).transpose(0, 2, 1) return p payoffs = np.array([[9, 8, 15], [10, 4, 15], [2, 30, 12]]) rng = np.random.default_rng() possible_path = np.array(list(itertools.combinations_with_replacement(range(0, 3), 100))) current_prob = rng.dirichlet(np.ones(3), size=1000) current_prob = current_prob.T p = rng.uniform(0., 1.0, size=1000) b = rng.uniform(0., 1.0, size=1000) c = rng.uniform(0., 1.0, size=1000) d = rng.uniform(0., 1.0, size=1000) all_payoffs = [] past_weight = np.zeros((1000,)) past_belief = np.zeros((3, 1000)) for t in range(0, possible_path.shape[1]): action = possible_path[:, t] num_path = np.arange(action.shape[0]) p_l1 = np.zeros((action.shape[0], 3)) p_l1[num_path, action] = 1.0 if t == 0: current_signal = current_prob.T * (1 - p)[:, None] + p_l1[:, None] * p[:, None] else: current_signal = p_l1[:, None] * p[:, None] + current_prob.transpose(1, 0, 2) * (1 - p)[:, None] current_weight = past_weight * c if t == 0: current_belief = (current_signal.transpose(1, 2, 0) + (past_belief * current_weight)[:, None].transpose(2, 0, 1)) / (1 + current_weight)[:, None][:, None] else: current_belief = (current_signal.transpose(1, 2, 0) + (past_belief * current_weight[:, None][:, None])) / (1 + current_weight)[:, None][:, None] current_payoffs = payoffs @ current_belief current_prob = get_p(current_payoffs) past_belief = current_belief past_weight = current_weight + 1
优化建议
1. 高效生成one-hot数组p_l1
原代码通过np.zeros初始化再赋值的方式效率较低,改用np.eye直接索引生成one-hot数组,利用NumPy内部优化减少内存操作:
# 替换原p_l1生成代码 p_l1 = np.eye(3)[action] # shape: (num_paths, 3)
2. 简化get_p函数的掩码逻辑
合并多次掩码操作,减少数组读写次数,同时用keepdims=True替代手动维度扩展,代码更简洁:
def get_p(payoffs): allmax = payoffs.max(axis=1, keepdims=True) findmax = payoffs - allmax # 批量生成所有需要置-1的掩码 mask_12 = (findmax[:, 1, :] == 0) & (findmax[:, 2, :] == 0) mask_01 = (findmax[:, 0, :] == 0) & (findmax[:, 1, :] == 0) mask_all = (findmax[:, 0, :] == 0) & (findmax[:, 1, :] == 0) & (findmax[:, 2, :] == 0) # 一次性应用掩码 findmax[:, 1, :] = np.where(mask_12 | mask_all, -1, findmax[:, 1, :]) findmax[:, 0, :] = np.where(mask_01 | mask_all, -1, findmax[:, 0, :]) p = np.where(findmax < 0, 0.0, 1.0).transpose(0, 2, 1) return p
3. 消除循环内的条件分支
通过初始化时调整变量维度,合并current_signal和current_belief的分支逻辑,减少循环内的条件判断开销:
# 初始化时调整current_prob维度,避免后续转置 current_prob = rng.dirichlet(np.ones(3), size=1000) # shape: (1000, 3) # 去掉原代码中的current_prob = current_prob.T # 初始化past_belief为(1,3,1000),统一后续计算维度 past_belief = np.zeros((1, 3, 1000)) past_weight = np.zeros((1000,)) # 循环内替换current_signal逻辑 if t == 0: current_signal = current_prob[None, :, :] * (1 - p)[None, :, None] + p_l1[:, None, :] * p[None, :, None] else: current_signal = p_l1[:, None, :] * p[None, :, None] + current_prob * (1 - p)[None, :, None] # 合并current_belief计算,消除分支 denominator = (1 + current_weight)[None, None, :] current_belief = (current_signal.transpose(1, 2, 0) + past_belief * current_weight[None, None, :]) / denominator
4. 利用广播替代显式维度扩展
将(1 + current_weight)[:, None][:, None]简化为(1 + current_weight)[None, None, :],借助NumPy广播机制自动匹配维度,减少数组复制:
denominator = (1 + current_weight)[None, None, :] current_belief = (current_signal.transpose(1, 2, 0) + past_belief * current_weight[None, None, :]) / denominator
5. 预分配结果数组(若需存储all_payoffs)
如果需要保存每次循环的current_payoffs,预先分配固定大小的数组,避免动态列表追加的开销:
# 初始化时预分配 all_payoffs = np.zeros((possible_path.shape[1],) + payoffs.shape[:1] + past_belief.shape[1:]) # 循环内直接赋值 all_payoffs[t] = current_payoffs
内容的提问来源于stack exchange,提问作者jasmine
相关产品推荐
相关产品推荐

