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

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

优化思路解析

  1. 扁平化空间维度:将每个模式的23×23空间维度压缩为一维,方便后续批量检查
  2. 批量条件检查:用any(dim=1)批量判断每个模式是否包含目标值,替代循环逐个检查
  3. 掩码合并:通过逻辑与操作得到最终保留的模式掩码,直接用于张量索引过滤

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.25 23:42:31