PyTorch如何高效生成排除指定值的二维序列张量
问题说明
现有存储取值范围为0到n-1整数的PyTorch张量,需要构造列数为n-1的张量,每一行对应原张量同位置的数值,内容为0到n-1的完整序列排除该位置数值后的结果。常规场景下原张量元素个数远大于n(即a.numel() >> n),且方案需要支持便捷扩展带额外批次维度的场景。
示例输入输出:
import torch n = 3 a = torch.Tensor([0, 1, 2, 1, 2, 0]).long() # 期望输出 b = [ [1, 2], [0, 2], [0, 1], [0, 2], [0, 1], [1, 2] ]
实现思路
因为n通常远小于张量a的元素总数,优先用预构造全量序列+掩码筛选的方案,避免循环带来的性能损耗:
- 先构造形状和
a对齐、最后一维长度为n的全量0~n-1序列 - 生成布尔掩码,标记每个位置上不等于原张量
a对应值的位置(即需要保留的元素位置) - 用掩码过滤掉被排除的元素,最终reshape得到最后一维长度为
n-1的目标结果
代码实现
def exclude_self_value(a: torch.Tensor, n: int) -> torch.Tensor: # 保证输入索引为整型,避免类型不匹配问题 a = a.long() # 构造全量0~n-1序列,自动对齐a的所有前置维度,最后一维长度为n full_seq = torch.arange(n, device=a.device, dtype=a.dtype).expand(*a.shape, n) # 生成保留掩码:最后一维上不等于对应a位置值的位置保留 mask = full_seq != a.unsqueeze(-1) # 过滤后调整形状到目标维度 return full_seq[mask].reshape(*a.shape, n-1)
测试示例效果:
n = 3 a = torch.tensor([0, 1, 2, 1, 2, 0]) b = exclude_self_value(a, n) print(b) # 输出与期望完全一致: # tensor([[1, 2], # [0, 2], # [0, 1], # [0, 2], # [0, 1], # [1, 2]])
批次维度扩展说明
该实现原生支持任意多的前置批次维度,无需额外修改:
- 若输入
a形状为(batch_size, seq_len),输出形状自动为(batch_size, seq_len, n-1) - 若输入
a形状为(B1, B2, L),输出形状自动匹配为(B1, B2, L, n-1)
所有运算全程和输入张量在同一设备上执行,无CPU/GPU数据拷贝开销,在a.numel() >> n的场景下性能远高于逐行循环、列表推导等实现。
性能优化提示:当n为固定常量时,可以预先缓存
torch.arange(n)的常量张量,进一步减少重复构造张量的开销。
内容的提问来源于stack exchange,提问作者Nagabhushan S N
相关产品推荐
相关产品推荐

