PyTorch中Tensor张量的计算法则
2026/9/6 2:05:12 网站建设 项目流程

一、索引

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=0dim=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来精准调整。😊

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询