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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.28 07:42:15