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

如何在PyTorch中通过匹配列表值筛选二维张量元素?

解决PyTorch二维张量筛选匹配指定列表元素的问题

首先,你遇到的报错大概率是因为没有将pandas Index类型的a转换为PyTorch张量,或是对torch.where的使用场景理解有误。下面给出正确的实现步骤:

步骤1:准备数据并转换类型

首先将pandas Index对象a转为PyTorch张量,因为torch.isin只接受张量作为匹配目标:

import torch
from pandas import Index

# 原二维张量
example = torch.tensor([[61078, 51477, 28492, 4290, 86920,  2216],
                        [26799, 76684, 23785, 18202, 14552, 98301]])
# 指定匹配列表
a = Index([61078, 23785, 2216], dtype='int64', length=3)
# 转换为PyTorch张量
a_tensor = torch.tensor(a.values)

步骤2:生成匹配掩码

用torch.isin生成布尔掩码,标记张量中哪些元素属于指定列表:

mask = torch.isin(example, a_tensor)
# 掩码结果:
# tensor([[ True, False, False, False, False,  True],
#         [False, False,  True, False, False, False]])

步骤3:提取匹配元素

方式1:提取所有匹配元素(一维张量)

直接用掩码索引原张量即可:

all_matched = example[mask]
print(all_matched)
# 输出:tensor([61078, 2216, 23785])

方式2:按行提取匹配元素(每行保留自身匹配项)

如果需要按行保留匹配元素,可遍历每行和对应的掩码:

row_matched = [row[row_mask] for row, row_mask in zip(example, mask)]
print(row_matched)
# 输出:[tensor([61078, 2216]), tensor([23785])]

方式3:统一每行长度(填充到指定长度)

如果需要让每行输出长度一致,可使用pad_sequence填充:

from torch.nn.utils.rnn import pad_sequence

padded_result = pad_sequence(row_matched, batch_first=True, padding_value=0)
print(padded_result)
# 输出:
# tensor([[61078, 2216,     0],
#         [23785,     0,     0]])

关于你提供的示例结果的说明

你给出的示例结果中包含大量不在指定列表a中的元素(比如28492、26799等),这和“筛选匹配元素”的需求不符,大概率是示例存在笔误。如果你的真实需求是其他逻辑(比如保留特定列的元素),可以补充细节后调整实现。

常见报错原因及解决

  1. 类型不兼容报错:因为a是pandas Index对象,直接传给torch.isin会报错,解决方法就是将其转为PyTorch张量(如步骤1所示)。
  2. torch.where使用错误:torch.where的作用是根据条件选择两个张量中的元素,而非直接提取匹配项。如果想用torch.where标记匹配元素,可以这样写:
# 将匹配元素保留,不匹配元素替换为0
marked_result = torch.where(mask, example, torch.zeros_like(example))
print(marked_result)
# 输出:
# tensor([[61078,     0,     0,     0,     0,  2216],
#         [    0,     0, 23785,     0,     0,     0]])

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.23 08:25:57