如何在PyTorch中高效从LongTensor集合移除指定LongTensor?
嘿,这个问题我之前也踩过坑!你碰到的KeyError其实是PyTorch Tensor和Python集合的特性冲突导致的,我来给你拆解下原因和靠谱的解决办法:
为什么会报错?
Python的集合(set)是靠哈希值来判断元素是否存在的,而PyTorch的Tensor是可变对象——哪怕两个Tensor里存的值完全一样,它们也是不同的Python对象,哈希值自然不一样。所以你定义的p和集合里的torch.LongTensor([1,3])看起来内容相同,但集合认不出它们是同一个元素,就抛出KeyError了。
高效解决办法
1. 转成不可变类型存集合(最推荐)
既然Tensor本身没法当可靠的集合元素,我们可以把它转成可哈希、不可变的tuple来存储,这样就能正常用集合的remove、in等操作了,效率还很高:
import torch # 把每个Tensor转成tuple后存入集合 ks = {tuple(torch.LongTensor([1, 3])), tuple(torch.LongTensor([2, 3])), tuple(torch.LongTensor([3, 3]))} p = torch.LongTensor([1, 3]) # 移除时也把目标Tensor转成tuple ks.remove(tuple(p)) print(ks) # 输出: {(2, 3), (3, 3)}
这个方法的转换开销极小,而且集合的查找/移除操作是O(1)的,适合绝大多数场景。
2. 用列表替代集合,按值匹配移除
如果你不想改存储方式,也可以把集合换成列表,遍历找到值匹配的元素再删除。注意这个方法会移除第一个匹配的元素,时间复杂度是O(n),适合元素不多的情况:
import torch ks = [torch.LongTensor([1, 3]), torch.LongTensor([2, 3]), torch.LongTensor([3, 3])] p = torch.LongTensor([1, 3]) # 遍历找到第一个值完全匹配的元素并删除 for idx, tensor in enumerate(ks): if torch.all(tensor == p): del ks[idx] break print(ks) # 输出: [tensor([2, 3]), tensor([3, 3])]
3. 自定义可哈希的Tensor包装类(进阶)
如果必须直接存储Tensor对象,可以写个简单的包装类,重写__hash__和__eq__方法,让它基于Tensor的值来判断相等和计算哈希:
import torch class HashableTensor: def __init__(self, tensor): self.tensor = tensor # 提前计算哈希值,避免重复计算开销 self._hash = hash(tuple(tensor.cpu().numpy())) def __hash__(self): return self._hash def __eq__(self, other): if not isinstance(other, HashableTensor): return False return torch.all(self.tensor == other.tensor) # 用包装类存储Tensor ks = {HashableTensor(torch.LongTensor([1, 3])), HashableTensor(torch.LongTensor([2, 3])), HashableTensor(torch.LongTensor([3, 3]))} p = HashableTensor(torch.LongTensor([1, 3])) ks.remove(p) # 可以通过.tensor属性拿到原Tensor print([ht.tensor for ht in ks]) # 输出: [tensor([2, 3]), tensor([3, 3])]
这个方法适合需要直接操作原Tensor的场景,但会有一点包装开销,按需选择就好。
内容的提问来源于stack exchange,提问作者Jesujoba Oluwadara ALABI
相关产品推荐
相关产品推荐

