第 3 章 索引与切片
学习目标
- 掌握一维张量的索引、切片与步长
- 掌握多维张量的
a[i, j]式索引与分维度切片 - 理解切片返回视图,修改视图会影响原张量
- 掌握
clone()复制、布尔索引与花式索引 - 知道 PyTorch 切片与 NumPy 的一个关键差异:不支持负步长
3.1 一维张量的索引
索引(indexing)是用位置取出元素。Python 的索引从 0 开始,支持负数从尾部数:
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():
print(a[0].item())输出:
103.2 一维张量的切片
切片(slicing)用 起点:终点:步长 取一段,含起点、不含终点:
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:
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):它不复制数据,而是「借用」原张量的一部分内存。修改视图,原张量同步变化:
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():
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 列」:
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)每个维度都可以独立切片,逗号分隔,: 表示该维全部:
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 的位置。先做比较运算得到布尔张量:
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(条件, 为真时的值, 为假时的值) 做逐元素选择:
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)用一个整数张量当索引,按位置收集元素:
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 3IndexError 是索引错误,信息 index 5 is out of bounds for dimension 0 with size 3 表示:在长度 3 的第 0 维上访问索引 5,越界了。先打印 shape 确认每个维度的长度,再检查索引值是否落在 [0, 长度-1] 内。
动手实践
- 创建
torch.arange(1, 11),用切片取出[4, 5, 6]、[2, 4, 6, 8, 10]、[1, 2, 3](三种写法)。 - 用
torch.flip把[1, 2, 3, 4, 5]倒序,再尝试[::-1]观察报错并复述原因。 - 创建 4×4 的单位矩阵
torch.eye(4),用切片取出主对角线(提示:逐行取m[i, i],或用第 2 章提示的torch.diagonal)。 - 对成绩张量
[85, 92, 78, 90, 88],用布尔索引取出所有 ≥ 90 的成绩。 - 验证「切片是视图」:对
torch.arange(6).reshape(2, 3)取[:, 1:],修改它,观察原张量变化。
常见错误
| 错误写法 | 现象 | 原因 |
|---|---|---|
a[::-1] | ValueError: step must be greater than zero | PyTorch 不支持负步长切片,用 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 再取值 |
| 修改花式索引结果,以为会影响原张量 | 原张量不变 | 花式索引返回副本,不是视图 |
章末练习
基础
- 创建
torch.arange(0, 20, 2),用索引取出第 3 个元素、最后一个元素、倒数第 2 个元素。 - 用切片取出
[10, 30, 50](从torch.arange(10, 60, 10)),再用torch.flip倒序整个张量。 - 创建 3×3 张量
[[1,2,3],[4,5,6],[7,8,9]],取出第 2 行、第 1 行第 2 列、最后一列。
提高
- 解释「切片是视图」并用一段 3 行代码演示:先创建张量、切片、改切片,再打印原张量。
- 用布尔索引从
torch.tensor([5, 12, 7, 18, 3, 20])中取出所有大于 10 且小于 20 的数(提示:用&连接两个比较条件,注意加括号)。
挑战
- 不运行代码,预测
torch.tensor([1, 2, 3, 4])[torch.tensor([True, False, True, False])]的结果,再运行验证。 - 用
torch.where把成绩[85, 92, 78, 90, 88]中 ≥ 90 的标记为 1、其余为 0,并与布尔索引写法对比;说明where的返回值形状。
章末自测
每题选择一个最佳答案。本书不附答案:完成后交由老师或 AI 老师批改讲解。
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])
- A.
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])
- A.
- 下列哪种方式得到的是视图(修改它会影响原张量)?
- A.
a.clone() - B.
a[1:3] - C.
torch.tensor(a) - D.
a[torch.tensor([0, 2])]
- A.
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. 返回空张量
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])
- A.
m[:, 1](m 为 3×3 张量)的形状是?- A.
torch.Size([3, 1]) - B.
torch.Size([3]) - C.
torch.Size([1, 3]) - D.
torch.Size([9])
- A.
- 把
[10, 20, 30, 40, 50]倒序,正确做法是?- A.
a[::-1] - B.
torch.flip(a, dims=[0]) - C.
a.reverse() - D.
a[-1:0:-1]
- A.
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)
- A.
- 布尔索引返回的是?
- A. 视图
- B. 副本
- C. 与花式索引相同都是视图
- D. 取决于张量大小
- 索引越界时会抛出?
- A.
ValueError - B.
RuntimeError - C.
IndexError - D.
TypeError
- A.
