如何在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
相关产品推荐
相关产品推荐

