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

优化PyTorch函数:用原生函数替代循环,将零元素替换为右侧首个非零值

优化PyTorch中“将零元素替换为右侧首个非零元素”的实现

需求:将张量中的零元素替换为其右侧首个非零元素,用PyTorch原生函数替代循环以提升性能。

原循环实现

last_val=0
for i in range(t.shape[0]-1, -1, -1):
    if (t[i] > 0):
        last_val = t[i]
    else:
        t[i] = last_val

输入与预期输出

输入:

t = torch.tensor([0, 0, 5, 0, 0, 7, 0, 8, 9])

预期输出:

tensor([5, 5, 5, 7, 7, 7, 8, 8, 9])

优化后的实现

def replace_zeros_with_right_nonzero(t):
    # 翻转张量,将“右侧首个非零”转化为“左侧首个非零”问题
    t_rev = t.flip(0)
    # 生成非零元素的掩码
    mask = t_rev != 0
    # 创建索引数组:非零位置保留自身索引,零位置设为-1
    idxs = torch.arange(t_rev.size(0), device=t.device)
    idxs = torch.where(mask, idxs, torch.tensor(-1, dtype=torch.long, device=t.device))
    # 累积取最大索引,得到每个位置最近的左侧非零元素索引
    last_nonzero_idxs = torch.cummax(idxs, dim=0)[0]
    # 根据索引填充值
    t_rev_filled = t_rev[last_nonzero_idxs]
    # 翻转回原顺序得到结果
    return t_rev_filled.flip(0)

# 测试示例
t = torch.tensor([0, 0, 5, 0, 0, 7, 0, 8, 9])
result = replace_zeros_with_right_nonzero(t)
print(result)  # 输出: tensor([5, 5, 5, 7, 7, 7, 8, 8, 9])

实现原理

  1. 张量翻转:将原问题中“找右侧首个非零元素”转换为“找左侧首个非零元素”,这样可以利用PyTorch的向量化累积操作处理,彻底避免Python循环的开销。
  2. 掩码与索引处理:通过掩码标记非零位置,将零位置的索引设为-1;利用torch.cummax累积获取每个位置最近的左侧非零元素索引——cummax会保留到当前位置为止的最大索引,零位置的-1不会覆盖之前的有效索引,正好符合“取最近左侧非零元素”的需求。
  3. 还原顺序:填充后的翻转张量再次翻转,即可得到原顺序下的目标结果。

性能优势

该实现完全基于PyTorch原生向量化操作,能充分利用CPU/GPU的并行计算能力。在处理大尺寸张量时,性能远优于逐元素循环的实现——循环会引入大量Python解释器开销,而向量化操作则由底层优化的C++/CUDA代码执行,效率提升显著。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.08 02:17:13