Skip to content

第 3 章 索引与切片

学习目标

  • 掌握一维张量的索引、切片与步长
  • 掌握多维张量的 a[i, j] 式索引与分维度切片
  • 理解切片返回视图,修改视图会影响原张量
  • 掌握 clone() 复制、布尔索引与花式索引
  • 知道 PyTorch 切片与 NumPy 的一个关键差异:不支持负步长

3.1 一维张量的索引

索引(indexing)是用位置取出元素。Python 的索引从 0 开始,支持负数从尾部数:

python
import torch

a = torch.tensor([10, 20, 30, 40, 50])
print(a[0])     # 第一个元素
print(a[-1])    # 最后一个元素

输出:

tensor(10)
tensor(50)

注意 a[0] 返回的是一个 0 维张量,打印为 tensor(10),不是 Python 整数 10。需要纯数字时用 .item():

python
print(a[0].item())

输出:

10

3.2 一维张量的切片

切片(slicing)起点:终点:步长 取一段,含起点、不含终点:

python
import torch

a = torch.tensor([10, 20, 30, 40, 50])
print(a[1:4])    # 索引 1、2、3
print(a[::2])    # 从 0 开始,隔一个取一个

输出:

tensor([20, 30, 40])
tensor([10, 30, 50])

PyTorch 不支持负步长:a[::-1] 会直接报错。

>>> import torch
>>> torch.tensor([10, 20, 30, 40, 50])[::-1]
Traceback (most recent call last):
  File "<stdin>", line 1, in <module>
ValueError: step must be greater than zero

倒序要用 torch.flip:

python
import torch

a = torch.tensor([10, 20, 30, 40, 50])
print(torch.flip(a, dims=[0]))

输出:

tensor([50, 40, 30, 20, 10])

dims=[0] 表示沿第 0 维(第一维)翻转。多维倒序在第 5 章形状操作里会再遇到。

3.3 视图与副本

切片返回的是视图(view):它不复制数据,而是「借用」原张量的一部分内存。修改视图,原张量同步变化:

python
import torch

a = torch.tensor([10, 20, 30, 40, 50])
view = a[1:4]
view[0] = 99
print(a)

输出:

tensor([10, 99, 30, 40, 50])

只改 view[0],但 a[1] 变成了 99——它们共享同一块内存。

需要独立副本时用 clone():

python
import torch

b = torch.tensor([10, 20, 30, 40, 50])
cop = b.clone()
cop[0] = 0
print("b:", b)
print("cop:", cop)

输出:

b: tensor([10, 20, 30, 40, 50])
cop: tensor([ 0, 20, 30, 40, 50])

改副本不影响原张量。记忆口诀:切片是视图(共享),clone() 是副本(独立)。视图带来的隐患在训练循环里很常见:无意中通过视图修改了参数张量。

3.4 多维张量的索引与切片

多维张量用逗号分隔每个维度的索引,m[i, j] 等价于「第 i 行第 j 列」:

python
import torch

m = torch.tensor([[1, 2, 3],
                  [4, 5, 6],
                  [7, 8, 9]])
print(m[0])        # 第 0 行
print(m[1, 2])     # 第 1 行第 2 列

输出:

tensor([1, 2, 3])
tensor(6)

每个维度都可以独立切片,逗号分隔,: 表示该维全部:

python
print(m[:, 1])      # 所有行的第 1 列
print(m[0:2, 1:])   # 前两行、从第 1 列开始

输出:

tensor([2, 5, 8])
tensor([[2, 3],
        [5, 6]])

m[:, 1] 取的是列,但结果是一维张量(形状 torch.Size([3])),不是 (3, 1) 的列向量。要保持二维形状,需要加 None 或用 reshape(第 5 章)。

3.5 布尔索引与花式索引

布尔索引用形状相同的布尔张量当「面具」,只取出为 True 的位置。先做比较运算得到布尔张量:

python
import torch

c = torch.tensor([10, 20, 30, 40, 50])
mask = c % 20 == 0
print(mask)
print(c[mask])

输出:

tensor([False,  True, False,  True, False])
tensor([20, 40])

布尔张量常用 torch.where(条件, 为真时的值, 为假时的值) 做逐元素选择:

python
import torch

c = torch.tensor([10, 20, 30, 40, 50])
print(torch.where(c % 20 == 0, torch.tensor(1), torch.tensor(0)))

输出:

tensor([0, 1, 0, 1, 0])

花式索引(fancy indexing)用一个整数张量当索引,按位置收集元素:

python
import torch

c = torch.tensor([10, 20, 30, 40, 50])
idx = torch.tensor([0, 2])
print(c[idx])

输出:

tensor([10, 30])

花式索引返回的是副本,不是视图——修改结果不会影响原张量。

3.6 读懂索引越界

索引越界是新手最常见的报错之一:

>>> import torch
>>> torch.tensor([10, 20, 30])[torch.tensor([0, 5])]
Traceback (most recent call last):
  File "<stdin>", line 1, in <module>
IndexError: index 5 is out of bounds for dimension 0 with size 3

IndexError 是索引错误,信息 index 5 is out of bounds for dimension 0 with size 3 表示:在长度 3 的第 0 维上访问索引 5,越界了。先打印 shape 确认每个维度的长度,再检查索引值是否落在 [0, 长度-1] 内。

动手实践

  1. 创建 torch.arange(1, 11),用切片取出 [4, 5, 6][2, 4, 6, 8, 10][1, 2, 3](三种写法)。
  2. torch.flip[1, 2, 3, 4, 5] 倒序,再尝试 [::-1] 观察报错并复述原因。
  3. 创建 4×4 的单位矩阵 torch.eye(4),用切片取出主对角线(提示:逐行取 m[i, i],或用第 2 章提示的 torch.diagonal)。
  4. 对成绩张量 [85, 92, 78, 90, 88],用布尔索引取出所有 ≥ 90 的成绩。
  5. 验证「切片是视图」:对 torch.arange(6).reshape(2, 3)[:, 1:],修改它,观察原张量变化。

常见错误

错误写法现象原因
a[::-1]ValueError: step must be greater than zeroPyTorch 不支持负步长切片,用 torch.flip
a[0] 当整数用,如 a[0] + 1 后再拼字符串TypeError 或打印成 tensor(10)0 维索引返回张量,要用 .item()
修改切片后以为不影响原张量原张量跟着变了切片是视图,共享内存;独立副本用 clone()
m[:, 1] 期望是列向量形状是 torch.Size([3])单列索引会降维,保持二维需 m[:, 1:2] 或加 None
c[torch.tensor([0, 5])]IndexError: index 5 is out of bounds ...索引越界;先查 shape 再取值
修改花式索引结果,以为会影响原张量原张量不变花式索引返回副本,不是视图

章末练习

基础

  1. 创建 torch.arange(0, 20, 2),用索引取出第 3 个元素、最后一个元素、倒数第 2 个元素。
  2. 用切片取出 [10, 30, 50](从 torch.arange(10, 60, 10)),再用 torch.flip 倒序整个张量。
  3. 创建 3×3 张量 [[1,2,3],[4,5,6],[7,8,9]],取出第 2 行、第 1 行第 2 列、最后一列。

提高

  1. 解释「切片是视图」并用一段 3 行代码演示:先创建张量、切片、改切片,再打印原张量。
  2. 用布尔索引从 torch.tensor([5, 12, 7, 18, 3, 20]) 中取出所有大于 10 且小于 20 的数(提示:用 & 连接两个比较条件,注意加括号)。

挑战

  1. 不运行代码,预测 torch.tensor([1, 2, 3, 4])[torch.tensor([True, False, True, False])] 的结果,再运行验证。
  2. torch.where 把成绩 [85, 92, 78, 90, 88] 中 ≥ 90 的标记为 1、其余为 0,并与布尔索引写法对比;说明 where 的返回值形状。

章末自测

每题选择一个最佳答案。本书不附答案:完成后交由老师或 AI 老师批改讲解。

  1. torch.arange(10)[1:4] 的结果是?
    • A. tensor([1, 2, 3, 4])
    • B. tensor([1, 2, 3])
    • C. tensor([0, 1, 2])
    • D. tensor([2, 3, 4])
  2. torch.tensor([10, 20, 30, 40, 50])[::2] 的结果是?
    • A. tensor([20, 40])
    • B. tensor([10, 30, 50])
    • C. tensor([10, 20, 30])
    • D. tensor([50, 40, 30])
  3. 下列哪种方式得到的是视图(修改它会影响原张量)?
    • A. a.clone()
    • B. a[1:3]
    • C. torch.tensor(a)
    • D. a[torch.tensor([0, 2])]
  4. torch.tensor([10, 20, 30, 40, 50])[::-1] 会?
    • A. 返回倒序张量
    • B. 报 ValueError: step must be greater than zero
    • C. 返回 tensor([50, 40, 30, 20, 10]) 的视图
    • D. 返回空张量
  5. m = torch.tensor([[1,2,3],[4,5,6],[7,8,9]]); m[1, 2] 的结果是?
    • A. tensor([4, 5, 6])
    • B. tensor(6)
    • C. 6
    • D. tensor([6])
  6. m[:, 1](m 为 3×3 张量)的形状是?
    • A. torch.Size([3, 1])
    • B. torch.Size([3])
    • C. torch.Size([1, 3])
    • D. torch.Size([9])
  7. [10, 20, 30, 40, 50] 倒序,正确做法是?
    • A. a[::-1]
    • B. torch.flip(a, dims=[0])
    • C. a.reverse()
    • D. a[-1:0:-1]
  8. c = torch.tensor([10, 20, 30]); c[c > 15] 的结果是?
    • A. tensor([20, 30])
    • B. tensor([10, 20, 30])
    • C. tensor([False, True, True])
    • D. tensor(20)
  9. 布尔索引返回的是?
    • A. 视图
    • B. 副本
    • C. 与花式索引相同都是视图
    • D. 取决于张量大小
  10. 索引越界时会抛出?
    • A. ValueError
    • B. RuntimeError
    • C. IndexError
    • D. TypeError