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

如何在PyTorch中向量化实现双层循环的张量替换逻辑

PyTorch向量化实现特定元素替换逻辑

问题描述

给定参数和张量如下:

import torch

vocab_size = 20
batch_size = 2
input_len = 5
output_len = 10
input_ids = torch.randint(0, vocab_size, (batch_size, input_len))
output_ids = torch.randint(0, vocab_size, (batch_size, output_len))

示例输入:

input_ids = tensor([[ 0,  8,  7, 12,  8],
                    [14, 15,  9,  7, 10]])
output_ids = tensor([[ 2, 8,  3, 15,  2, 19,  7,  1, 19,  8],
                     [10,  8,  0,  7, 16,  0,  6,  2, 16, 13]])

需要对output_ids执行如下替换规则:

  • 若元素存在于对应batch的input_ids中,且不属于{0,1,2}
  • 则替换为 vocab_size + 该元素在input_ids中首次出现的索引

现有双层循环实现(可正常运行但效率较低):

for i in range(batch_size):
    for k, value in enumerate(output_ids[i]):
        if value in input_ids[i] and value not in [0, 1, 2]:
            output_ids[i][k] = vocab_size + torch.where(input_ids[i] == value)[0][0]

示例输出:

tensor([[ 2, 21,  3, 15,  2, 19, 22,  1, 19, 24],
        [24,  8,  0, 23, 16,  0,  6,  2, 16, 13]])

向量化实现方案

通过PyTorch的张量广播、scatter映射和掩码操作,可完全替代循环实现该逻辑,且性能更优:

1. 构建input元素到首次索引的映射

为每个batch创建映射张量,记录每个元素在input_ids中的首次出现索引:

# 初始化映射张量,未出现的元素标记为-1
input_to_idx = torch.full((batch_size, vocab_size), -1, dtype=torch.long, device=input_ids.device)
# 生成batch维度的索引
batch_indices = torch.arange(batch_size).unsqueeze(1)
# scatter操作:将input_ids中每个元素的首次索引写入映射(重复元素不会覆盖已写入的首次索引)
input_to_idx.scatter_(1, input_ids, torch.arange(input_len, device=input_ids.device).repeat(batch_size, 1))

此时input_to_idx[i][v]表示第i个batch中元素v的首次出现索引,若v不在input_ids[i]中则为-1。

2. 生成替换值与匹配掩码

计算每个output元素的候选替换值,并筛选符合条件的元素:

# 获取每个output元素对应的替换值:vocab_size + 首次索引
replace_values = vocab_size + input_to_idx[batch_indices.repeat(1, output_len), output_ids]
# 构建掩码:元素不在{0,1,2} 且 在input中存在(映射索引不为-1)
mask = (~output_ids.isin(torch.tensor([0,1,2], device=output_ids.device))) & (input_to_idx[batch_indices.repeat(1, output_len), output_ids] != -1)

3. 应用替换逻辑

用torch.where根据掩码完成替换:

output_ids_vect = torch.where(mask, replace_values, output_ids)

完整可运行代码

import torch

vocab_size = 20
batch_size = 2
input_len = 5
output_len = 10

# 示例输入张量
input_ids = torch.tensor([[ 0,  8,  7, 12,  8],
                          [14, 15,  9,  7, 10]])
output_ids = torch.tensor([[ 2, 8,  3, 15,  2, 19,  7,  1, 19,  8],
                           [10,  8,  0,  7, 16,  0,  6,  2, 16, 13]])

# 步骤1:构建元素到首次索引的映射
input_to_idx = torch.full((batch_size, vocab_size), -1, dtype=torch.long, device=input_ids.device)
batch_indices = torch.arange(batch_size).unsqueeze(1)
input_to_idx.scatter_(1, input_ids, torch.arange(input_len, device=input_ids.device).repeat(batch_size, 1))

# 步骤2:生成替换值和掩码
replace_values = vocab_size + input_to_idx[batch_indices.repeat(1, output_len), output_ids]
mask = (~output_ids.isin(torch.tensor([0,1,2], device=output_ids.device))) & (input_to_idx[batch_indices.repeat(1, output_len), output_ids] != -1)

# 步骤3:执行替换
output_ids_vect = torch.where(mask, replace_values, output_ids)

print(output_ids_vect)

运行结果与循环实现完全一致:

tensor([[ 2, 21,  3, 15,  2, 19, 22,  1, 19, 24],
        [24,  8,  0, 23, 16,  0,  6,  2, 16, 13]])

方案优势

  • 无循环设计,利用PyTorch向量化操作大幅提升效率,在大batch或长序列场景下效果显著
  • 全张量操作,自动支持GPU加速,适配大规模训练场景

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.29 12:34:53