PyTorch中索引切片是否为视图?是否有代理式视图模块?
PyTorch切片视图与索引代理相关问题解答
一、你的理解完全正确
- 常规连续切片(比如
b = a[3:10, 2:5])属于PyTorch基础切片操作,会创建原张量的View(视图)——不复制数据,新张量只是原张量内存的别名,修改视图内容会直接影响原张量。 - 你提到的索引切片(注意你写的
a[[1,2,3] : [5,11]]是无效语法,正确高级索引写法比如a[[1,2,3], :]或a[:, [5,11]])属于高级索引,这类操作会触发数据复制,生成的新张量和原张量内存完全独立,修改新张量不会影响原张量。
二、模拟索引视图的通用实现
PyTorch原生没有专门提供这类存储索引作为代理的模块,但可以自己实现更通用的代理类,比你给出的IXView功能更完善:
class IndexedView: def __init__(self, base_tensor, indices): self.base = base_tensor self.indices = indices # 支持任意维度的索引张量或列表 def __getitem__(self, key): # 先处理索引的切片/子索引,再映射到原张量 if isinstance(key, tuple): idx_key, *remaining = key return self.base[self.indices[idx_key], *remaining] else: return self.base[self.indices[key]] # 转发原张量的常用属性和方法,避免重复实现 def __getattr__(self, name): return getattr(self.base, name)
使用示例
import torch # 创建原张量 a = torch.randn(10, 5) # 定义要代理的索引 target_indices = torch.tensor([1, 3, 5, 7]) # 创建代理视图 proxy_view = IndexedView(a, target_indices) # 各种访问方式都能映射到原张量的对应索引 print(proxy_view[0]) # 等价于 a[target_indices[0]] print(proxy_view[1:3]) # 等价于 a[target_indices[1:3]] print(proxy_view[:, 2]) # 等价于 a[target_indices, 2] print(proxy_view.mean(dim=1)) # 直接调用原张量的mean方法,计算索引部分的均值
这个类支持多维索引、切片操作,还能自动转发原张量的属性和方法,比基础版本更通用。
内容的提问来源于stack exchange,提问作者sten
相关产品推荐
相关产品推荐

