一、索引
1. 基本索引(取一个“点”)
直接用整数定位,结果会降维(去掉被取的那一维)。
x = torch.tensor([[1, 2, 3], [4, 5, 6]]) print(x[0, 1]) # 输出 2(第0行第1列,是个标量,0维) print(x[1]) # 输出 [4, 5, 6](第1行,变成1维)2. 切片索引(取一段“连续区间”)
使用start:stop:step,结果保持原维度不变(除非范围缩成1)。
x = torch.tensor([[1, 2, 3], [4, 5, 6]]) print(x[:, 0]) # 输出 [1, 4](就是你看过的那张图:所有行的第0列,变成1维) print(x[:, 0:1]) # 输出 [[1], [4]](注意!加了切片后,保持2维,形状为(2,1)) print(x[1:, :2]) # 输出 [[4, 5]](从第1行开始,取前两列)逗号,前后表示不同维度,冒号:表示全选
在print(x[:, 0:1])中,冒号:表示行全选,0:1代表列的索引是0:1,遵循左闭右开原则,只取第一列的元素
关键知识点:
因为在列的位置用的是切片:(哪怕是0:1这种只取一列的切片),它也会保留维度。所以取出来的结果依然是二维的(行还在,列也在),形状为(2, 1),写成矩阵就是两行一列。
而最开始的例子基本索引里是只取一个点x[:, 0](列位置是单个数字,没有冒号),那就是降维,结果会变成一维[1, 4]。有冒号就保留维度,没冒号(单数字)就降维。
3. 省略号(...):高维张量的“偷懒神器”
当你处理图片(B, C, H, W)时,不想写一堆冒号。...代表“所有未写出的中间维度”。
# 取所有批次、所有通道的左上角像素 (H=0, W=0) x[..., 0, 0] # 等价于 x[:, :, 0, 0]4. 花式索引(按指定列表“跳跃”取值)
传入列表或张量,结果形状由索引张量的形状决定。注意它是配对取值的。
x = torch.tensor([[1, 2], [3, 4], [5, 6]]) print(x[[0, 2], [0, 1]]) # 输出 [1, 6](取(0,0)和(2,1)两个点,配对!) print(x[[0, 2]]) # 输出 [[1, 2], [5, 6]](取第0行和第2行)核心规则:当在同一个方括号里写了两个列表(中间用逗号隔开),PyTorch 会像“拉拉链”一样,把两个列表按位置一对一配对取值。
第一个列表
[0, 2]:提供行号。第二个列表
[0, 1]:提供列号。
配对过程(逐对匹配):
取第1对:第一个列表的
0(行) + 第二个列表的0(列) → 取(行0, 列0)→ 数字1取第2对:第一个列表的
2(行) + 第二个列表的1(列) → 取(行2, 列1)→ 数字6
所以结果就是[1, 6]。
特别注意:它不是取“行0和行2”与“列0和列1”的全部交叉组合(即不是取
[[1,2],[5,6]]),而是严格的一对一配对。因为两个列表长度相同,所以结果是一维的。跟基础索引有点类似,不过基础索引是一个点,这里是多个点
5. 布尔索引(按条件“筛选”数据)
传入一个 True/False 的掩码(Mask),极其常用。结果会自动拉成一维。
x = torch.tensor([-1, 2, -3, 4]) mask = x > 0 print(x[mask]) # 输出 [2, 4](只把大于0的掏出来) # 高级用法:把小于0的变成0(类似ReLU效果) x[x < 0] = 0 print(x) # 输出 [0, 2, 0, 4]⚠️ 终极致命陷阱:视图(View) vs 副本(Copy)
这是所有新手都会踩的坑,一定要记住:
- 基本索引、切片、省略号:返回的是原数据的视图(View),共享内存。改了这个结果,原张量也会变!
y = x[:, 0] # 切片取第一列 y[0] = 999 print(x) # 原张量的对应位置也变成了 999 !花式索引、布尔索引:返回的是副本(Copy),不共享内存。改了没关系,原张量纹丝不动。
解决办法:如果你不希望改动原数据,在切片后加上.clone():y = x[:, 0].clone()。
一句话记忆口诀
逗号分维,冒号全选,花式配对,布尔筛选。切片改值会伤及原身,保险起见加个 clone。😊
二、换轴(主要关注x.T和x.transpose()的差别)
1. 在 PyTorch 中(你最关心的场景)
x.T:它只做一件事——交换第 0 轴和第 1 轴(即dim=0和dim=1)。对于二维矩阵,这是标准的转置;但对于高维张量,它忽略第 2、第 3 等后面的轴。torch.transpose(x, dim0, dim1):你必须手动指定要交换哪两个轴。它可以交换任意两个维度(比如第 1 轴和第 3 轴)。
用三维张量举例(形状(2, 3, 4)):
import torch x = torch.randn(2, 3, 4) # 1. 使用 x.T print(x.T.shape) # 输出:torch.Size([3, 2, 4]) # 原因:只交换了原来的第0轴(2)和第1轴(3),第2轴(4)原地不动。 # 2. 使用 transpose(0, 1) —— 这个结果和 x.T 一样 print(torch.transpose(x, 0, 1).shape) # 输出:torch.Size([3, 2, 4]) (一样) # 3. 使用 transpose(1, 2) —— 交换第1轴和第2轴 print(torch.transpose(x, 1, 2).shape) # 输出:torch.Size([2, 4, 3]) (完全不同!因为第2轴和第1轴互换了)所以结论:在 PyTorch 里,x.T永远只等于x.transpose(0, 1)。只要你想交换的不是 0 和 1,它们的结果就绝对不一样。
2. 在 NumPy 中
NumPy 的情况更复杂,因为transpose还有一个“空手套白狼”的用法:
x.T:同 PyTorch,只交换 0 和 1 轴。x.transpose()(不带参数):这是反转所有轴的顺序。比如形状(2, 3, 4)会变成(4, 3, 2),这跟x.T的(3, 2, 4)完全不是一回事。
3. 如果你想要“全部维度”转置怎么办?
在 PyTorch 中,如果你想一次性把所有维度完全颠倒过来(比如(2,3,4)->(4,3,2)),你不能用x.T,也不能用两两交换的transpose,而是要用x.permute():
# 把所有维度倒序排列 x.permute(2, 1, 0).shape # 输出 (4, 3, 2)🎯 终极记忆口诀(防止大脑短路)
x.T是“小跟班”,只敢动前两排(交换 0 和 1)。transpose是“指挥官”,想换哪两排就写哪两排(但要写参数)。想颠倒是非全部反过来?用
permute列出所有轴的新顺序。
实战建议:如果你的数据是二维表格(行列),用哪个都一样;如果你的数据是图片(Batch, Channel, Height, Width),千万别用x.T(它会把 Batch 和 Channel 搞混),一定要用transpose(1, 2)或permute来精准调整。😊