PyTorch函数优化:移除循环实现张量模式过滤
高效过滤PyTorch张量中存活模式的实现优化
问题说明
处理维度为torch.Size([51, 265, 23, 23])的张量:
- 第一维度:时间步
- 第二维度:独立模式
- 最后两个维度:模式的空间尺寸
规则:每个模式最多有[-1,0,1]三种状态,仅保留最后一个时间步中包含全部三种状态的“存活”模式,过滤掉“死亡”模式。
原始实现(存在性能瓶颈)
以下代码可正常运行,但因包含for循环,处理大张量时效率较低:
def filter_patterns(tensor_sims): # 获取需要保留的模式索引 keep_indices = torch.tensor([i for i in range(tensor_sims.shape[1]) if tensor_sims[-1,i].unique().numel() == 3]) # 过滤张量 tensor_sims = tensor_sims[:, keep_indices] print(f'Number of patterns: {tensor_sims.shape[1]}') return tensor_sims
优化后的矢量化实现
通过矢量化操作完全移除循环,运行速度大幅提升:
def filter_patterns(tensor_sims): # 扁平化最后一个时间步的空间维度 x_ = tensor_sims[-1].flatten(1) # 创建掩码,分别检查每个模式是否包含-1、0、1 mask_minus_one = (x_ == -1).any(dim=1) mask_zero = (x_ == 0).any(dim=1) mask_one = (x_ == 1).any(dim=1) # 合并掩码:同时满足三个条件的模式才保留 mask = mask_minus_one.logical_and(mask_zero).logical_and(mask_one) # 过滤张量 tensor_sims = tensor_sims[:, mask] print(f'Number of patterns: {tensor_sims.shape[1]}') return tensor_sims
优化思路解析
- 扁平化空间维度:将每个模式的23×23空间维度压缩为一维,方便后续批量检查
- 批量条件检查:用
any(dim=1)批量判断每个模式是否包含目标值,替代循环逐个检查 - 掩码合并:通过逻辑与操作得到最终保留的模式掩码,直接用于张量索引过滤
内容的提问来源于stack exchange,提问作者Fabrizio Brown
相关产品推荐
相关产品推荐

