如何按索引删除PyTorch三维tensor中的指定行
PyTorch三维张量删除指定行实现方法
你提到的「第2行」对应三维张量shape=(4,3,2)的**第0个维度(最外层维度)**的元素,删除后第0维长度从4变为3,符合你给出的预期结果。我们默认日常表述的「第n行」是1索引,PyTorch张量索引为0索引,因此删除第n行对应要排除的索引为n-1。
通用实现方法
下面给出两种最常用的实现方式,都可以适配任意维度张量删除指定位置元素的需求:
方法1:张量拼接(最简洁)
直接将要保留的前后两段张量拼接即可,代码如下:
import torch # 你的原张量 t = torch.tensor([[[0.4003, 0.2742], [0.9414, 0.1222], [0.9624, 0.3063]], [[0.9600, 0.5381], [0.5758, 0.8458], [0.6342, 0.5872]], [[0.5891, 0.9453], [0.8859, 0.6552], [0.5120, 0.5384]], [[0.3017, 0.9407], [0.4887, 0.8097], [0.9454, 0.6027]]]) n = 2 # 要删除的第n行(1索引) delete_idx = n - 1 # 转成PyTorch的0索引 # 在第0维拼接要保留的部分 new_t = torch.cat([t[:delete_idx], t[delete_idx+1:]], dim=0) print(new_t.shape) # 输出 torch.Size([3, 3, 2])
运行后得到的张量和你给出的预期结果完全一致。
方法2:掩码过滤(适配批量删除场景)
如果需要同时删除多行,用掩码的方式更灵活:
n = 2 delete_idx = n - 1 # 生成和第0维等长的掩码,初始为True(保留) mask = torch.ones(t.shape[0], dtype=torch.bool) # 将要删除的位置设为False mask[delete_idx] = False # 过滤得到新张量 new_t = t[mask]
其他维度的行删除说明
如果你后续需要删除其他维度的「行」,只需要调整对应参数即可:
- 要删除第1维(中间长度为3的维度)的第n行:把
dim参数改为1,掩码长度改为t.shape[1] - 要删除第2维(最内层长度为2的维度)的第n行:把
dim参数改为2,掩码长度改为t.shape[2]
内容的提问来源于stack exchange,提问作者Someone
相关产品推荐
相关产品推荐

