· 深度学习 ·阅读时长约 9 分钟

PyTorch 张量的索引与形状操作

从深度学习训练实战出发,系统讲清索引、形状变换与常见运算,帮助你减少维度错误并提升代码可维护性。

PyTorch张量索引形状操作

所属专题:PyTorch 深度学习基础 (pytorch·02)

读前提示(AI/机器学习视角)

  • 适合人群:已经会创建张量,准备进入“数据切片、批处理、维度变换”实操阶段的学习者。
  • 前置知识:建议先掌握张量创建、dtype/device 基本概念;熟悉二维数组行列语义。
  • 读完收获:你会知道如何稳定地做样本筛选、特征抽取、批维调整与张量重排,并能定位常见报错(如维度不匹配、view 连续性问题)。

@[toc]

1 张量索引

在训练代码里,索引几乎贯穿全流程:取一个 mini-batch、选择某些特征列、筛掉无效样本、提取指定通道。掌握索引语义,能显著减少“结果形状不符合预期”的调试时间。

1.1 简单行列索引和列表索引

import torch


# 1. 简单行列索引
def tensor_basic_row_col_indexing():

    # 固定随机数种子
    torch.manual_seed(0)

    data = torch.randint(0, 10, [4, 5])
    print(data)
    print('-' * 30)

    # 1.1 获得指定的某行元素
    # print(data[2])

    # 1.2 获得指定的某个列的元素
    # 逗号前面表示行, 逗号后面表示列

    # 冒号表示所有行或者所有列
    # print(data[:, :])

    # 表示获得第3列的元素
    print(data[:, 2])

    # 获得指定位置的某个元素
    print(data[1, 2], data[1][2])

    # 表示先获得前三行,然后再获得第三列的数据
    print(data[:3, 2])

    # 表示获得前三行的前两列
    print(data[:3, :2])


# 2. 列表索引
def tensor_advanced_list_indexing():


    # 固定随机数种子
    torch.manual_seed(0)

    data = torch.randint(0, 10, [4, 5])
    print(data)
    print('-' * 30)

    # 如果索引的行列都是一个1维的列表,那么两个列表的长度必须相等
    # print(data[[0, 1, 2], [2, 4]]) # 报错,索引位置都是一维,必须匹配
    # 解决方法:如果不想前后维数一样,就采用二维数组
    # 使用二维数组进行索引得到仍然为二维数组
    print(data[[[0],[1],[2]],[2,4]])


    # 1.表示获得 (0, 0)、(2, 1)、(3, 2) 三个位置的元素
    # 使用一维数组进行索引,得到的是一维
    print(data[[0, 2, 3], [0, 1, 2]])

    # 2。表示获得 0、2、3 行的 0、1、2 列
    # print(data[[[0], [2], [3]], [0, 1, 2]])

补充说明(常见易错点)

  • 当行索引和列索引都使用一维列表时,PyTorch 会按“位置配对”取值(类似坐标点采样),不是取笛卡尔积子矩阵。
  • 如果你想得到“多行 × 多列”的子矩阵,通常需要把其中一个索引改成二维(或分步索引),避免形状不符合预期。

1.2 布尔索引和多维索引

import torch


# 1. 布尔索引
def tensor_boolean_indexing_examples():

    torch.manual_seed(0)
    data = torch.randint(0, 10, [4, 5])
    print(data)

    # 1. 希望能够获得该张量中所有大于3的元素
    # 所有元素与3进行比较,大于返回True,小于返回False
    # 返回一个布尔类型的张量
    print(data > 3)

    # 对于张量中的所有元素进行筛选,变为一维的张量
    print(data[data > 3])


    # 2. 希望返回第2列元素大于6的行
    # 先获取到第二列数据,然后进行比较,得到布尔张量
    # 然后再进行行索引

    # 想要获取到行,在行索引的位置传入布尔张量
    print(data[:,1] > 6) # tensor([ True,  True, False, False])
    print(data[data[:, 1] > 6]) # 选择前两行

    # 3. 希望返回第2行元素大于3的所有列
    # 想要获取到列,在列的位置传入布尔索引
    print(data[:, data[1] > 3])


# 2. 多维索引
def tensor_multi_dim_indexing_examples():

    torch.manual_seed(0)
    data = torch.randint(0, 10, [3, 4, 5])
    print(data)
    print('-' * 30)

    # 按照第0个维度选择第0元素,4行5列元素
    print(data[0, :, :])
    print('-' * 30)

    # 按照第1个维度选择第0元素
    print(data[:, 0, :])
    print('-' * 30)

    # 按照第2个维度选择第0元素
    print(data[:, :, 0])
    print('-' * 30)

实战建议

  • 布尔索引非常适合做“样本过滤”(例如筛掉异常值、仅保留某类标签),但要注意:结果经常会被拉平成一维,需要后续显式 reshape
  • 多维张量索引时,建议在关键步骤打印 shape,把“每个维度代表什么”(如 batch, channel, height, width)写进注释,长期收益很高。

2 张量的形状操作

形状操作是深度学习代码最容易出 bug 的部分之一。网络层通常对输入维度有严格约束,因此“变形前后元素总数是否一致、维度顺序是否正确、内存是否连续”是三个关键检查点。

2.1 reshape函数

  • 保证张量元素个数不变的情况下改变张量的形状
  • 在神经网络中,不同层中的数据形状不同
import torch


def tensor_reshape_examples():

    torch.manual_seed(0)
    data = torch.randint(0, 10, [4, 5])

    # 查看张量的形状
    print(data.shape, data.shape[0], data.shape[1])
    # print(data.size(), data.size(0), data.size(1)) # 与上述方法结果一致

    # 修改张量的形状
    new_data = data.reshape(2, 10)
    print(new_data)

    # 注意: 转换之后的形状元素个数得等于原来张量的元素个数
    # 原来有多少个元素,转换之后就有多少个元素
    # new_data = data.reshape(1, 10)
    # print(new_data)

    # 使用-1代替省略的形状
    # 转换为指定的行数,列数指定为-1,可以进行自动匹配列数
    new_data = data.reshape(5, -1)
    print(new_data)

    # 转换为两列,自动进行计算行数
    new_data = data.reshape(-1, 2)
    print(new_data)

reshape 在多数场景下是首选:语义直观、可读性强。建议优先使用 -1 自动推断一个维度,减少手工计算出错概率。

2.2 transpose和permute函数的使用

  • reshape函数更改形状,reshape会重新计算张量的维度,有时候不需要重新计算张量的维度,只要调整张量维度的顺序即可,可以使用transpose函数和permute函数
  • transpose函数每次只能交换两个维度
  • permute函数可以一次交换多个维度
import torch


# 1. transpose 函数
def tensor_transpose_examples():

    torch.manual_seed(0)
    data = torch.randint(0, 10, [3, 4, 5])

    new_data = data.reshape(4, 3, 5)
    print(new_data.shape) # torch.Size([4, 3, 5])

    # 直接交换两个维度的值
    new_data = torch.transpose(data, 0, 1)
    print(new_data.shape) # torch.Size([4, 3, 5])

    # 缺点: 一次只能交换两个维度
    # 把数据的形状变成 (4, 5, 3)
    # 进行第一次交换: (4, 3, 5)
    # 进行第二次交换: (4, 5, 3)
    new_data = torch.transpose(data, 0, 1)
    new_data = torch.transpose(new_data, 1, 2)
    print(new_data.shape)


# 2. permute 函数
def tensor_permute_examples():

    torch.manual_seed(0)
    data = torch.randint(0, 10, [3, 4, 5])

    # permute 函数可以一次性交换多个维度
    new_data = torch.permute(data, [1, 2, 0])
    print(new_data.shape)

transpose/permute 只改变维度顺序,不改变元素总数。典型场景包括:图像张量在 HWCCHW 间转换、序列模型中交换时间维和 batch 维。

2.3 view和contiguous函数

  • view函数改变张量的形状,只能用于存储在整块内存中的张量,具有一定的局限性。
  • pytorch中有些张量是由不同的数据块组成,并没有存储在整块的内存中,view函数无法对于这种张量进行变形处理
  • 一个张量经过了transpose或者permute函数的处理之后,就无法使用view函数进行形状操作
  • 先用contiguous将非连续内存空间转换为连续内存空间,然后再使用view函数进行更改张量形状
import torch


# 1. view 函数的使用
def tensor_view_basic():

    data = torch.tensor([[10, 20, 30], [40, 50, 60]])
    data = data.view(3, 2)
    print(data.shape)

    # is_contiguous 函数来判断张量是否是连续内存空间(整块的内存)
    print(data.is_contiguous())


# 2. view 函数使用注意
def tensor_view_with_contiguous():

    # 当张量经过 transpose 或者 permute 函数之后,内存空间基本不连续
    # 此时,必须先把空间连续,才能够使用 view 函数进行张量形状操作

    data = torch.tensor([[10, 20, 30], [40, 50, 60]])
    print('是否连续:', data.is_contiguous())
    data = torch.transpose(data, 0, 1)
    print('是否连续:', data.is_contiguous())

    # 此时,在不连续内存的情况使用 view 会怎么样呢?
    data = data.contiguous().view(2, 3)
    print(data) 

为什么会报错?

  • view 要求底层内存连续;transpose/permute 后得到的张量通常是“非连续视图”。
  • 因此遇到 view 报错时,先检查 is_contiguous(),必要时先 contiguous()view()
  • 如果你更关注代码鲁棒性,可优先用 reshape,它会在需要时自动处理拷贝。

2.4 squeeze和unsqueeze函数用法

  • squeeze函数可以将维度为1的维度进行删除
  • unsqueeze函数给张量增加维度为1的维度
import torch


# 1. squeeze 函数使用
def tensor_squeeze_examples():

    # 四维张量
    data = torch.randint(0, 10, [1, 3, 1, 5])
    print(data.shape)

    # 维度压缩, 默认去掉所有的1的维度
    # squeeze() - 默认去掉所有维度为1的函数
    # squeeze(0) - 删除第一个位置的为1的维度
    # 传入维度的索引值
    new_data = data.squeeze(0)
    print(new_data.shape)  # torch.Size([3, 5])

    # 指定去掉某个1的维度
    new_data = data.squeeze(2)
    print(new_data.shape)


# 2. unsqueeze 函数使用
def tensor_unsqueeze_examples():

    data = torch.randint(0, 10, [3, 5])
    print(data.shape) # torch.Size([1, 3, 1, 5])


    # 可以在指定位置增加维度
    # -1 代表最后一个维度
    new_data = data.unsqueeze(-1)
    print(new_data.shape)

在训练中,unsqueeze 常用于给单样本补 batch 维,squeeze 常用于去掉冗余的长度为 1 的维度。
建议显式传入维度索引(如 squeeze(1)),避免误删掉不该删的维度。

2.5 张量更改形状小结

  1. reshape 函数可以在保证张量数据不变的前提下改变数据的维度.
  2. transpose 函数可以实现交换张量形状的指定维度, permute 可以一次交换更多的维度.
  3. view 函数也可以用于修改张量的形状, 但是它要求被转换的张量内存必须连续,所以一般配合 contiguous 函数使用.
  4. squeeze 和 unsqueeze 函数可以用来增加或者减少维度.

3 常见运算函数

这一节中的 mean/sum/log 等运算是损失计算、统计指标与特征标准化的基础操作。
实践中请重点关注:数据类型(整型/浮点型)、按哪个维度聚合(dim)、是否保留维度(可结合 keepdim=True)。

  • mean()
  • sum()
  • pow(n)
  • sqrt()
  • exp()
  • log() - 以e为底的对数
  • log2()
  • log10()
import torch


# 1. 均值
def tensor_mean_examples():

    torch.manual_seed(0)
    # data = torch.randint(0, 10, [2, 3], dtype=torch.float64)
    data = torch.randint(0, 10, [2, 3]).double()
    # print(data.dtype)

    print(data)
    # 默认对所有的数据计算均值
    print(data.mean())
    # 按指定的维度计算均值
    print(data.mean(dim=0)) # 竖向计算 按列计算
    print(data.mean(dim=1)) # 横向计算 按行计算


# 2. 求和
def tensor_sum_examples():

    torch.manual_seed(0)
    data = torch.randint(0, 10, [2, 3]).double()

    print(data.sum())
    print(data.sum(dim=0))
    print(data.sum(dim=1))


# 3. 平方
def tensor_pow_examples():

    torch.manual_seed(0)
    data = torch.randint(0, 10, [2, 3]).double()
    print(data)
    data = data.pow(2)
    print(data)


# 4. 平方根
def tensor_sqrt_examples():

    torch.manual_seed(0)
    data = torch.randint(0, 10, [2, 3]).double()
    print(data)
    data = data.sqrt()
    print(data)


# 5. e多少次方
def tensor_exp_examples():

    torch.manual_seed(0)
    data = torch.randint(0, 10, [2, 3]).double()
    print(data)
    data = data.exp()
    print(data)


# 6. 对数
def tensor_log_examples():

    torch.manual_seed(0)
    data = torch.randint(0, 10, [2, 3]).double()
    print(data)
    data = data.log()     # 以e为底
    data = data.log2()    # 以2为底
    data = data.log10()   # 以10为底
    print(data)

评论