Pytorch入门5-Pytorch张量的索引

Pytorch入门5-Pytorch张量的索引

张量的索引与numpy中的ndarray类似。

一、基础索引

示例代码如下:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
# 创建张量
tensor = torch.randint(1,10, (2,3,4))
print(tensor)
tensor([[[4, 8, 5, 5],
[4, 3, 8, 8],
[3, 3, 8, 4]],

[[8, 2, 3, 3],
[4, 9, 6, 7],
[4, 4, 2, 2]]])

# 获取第0个维度下标为1的元素
print(tensor[1])
tensor([[8, 2, 3, 3],
[4, 9, 6, 7],
[4, 4, 2, 2]])

# 获取第0个维度下标为1的元素,第1个维度下标为0的元素
print(tensor[1, 0])
tensor([8, 2, 3, 3])

# 获取第0个维度下标为1,第1个维度下标为0,第三个维度0-1的元素
print(tensor[1, 0, 0:2])
tensor([8, 2])

# 或
print(tensor[1, 0, [0,1]])
tensor([8, 2])

# 获取第0个维度下所有,第1个维度下0-1的元素,第2个维度为3的元素
print(tensor[:, 0:2, 3])
tensor([[5, 8],
[3, 7]])

通过以上代码,可以发现的规律为,中括号中对应位置的值,表示获取对应维度的元素。

二、列表(花式)索引

列表(花式)索引,表示在索引时,传入的是一个列表的值。下面分开来讲解。

1
2
3
4
5
6
7
8
9
10
11
12
t1 = torch.randint(1,10, (3,4))  
tensor([[4, 2, 5, 9],
[9, 1, 4, 9],
[1, 1, 7, 2]])

#两次获取0索引
## 从第0个维度获取元素,获取下标为0,0,1,2的元素
print(t1[[0,0,1,2]])
tensor([[4, 2, 5, 9],
[4, 2, 5, 9],
[9, 1, 4, 9],
[1, 1, 7, 2]])

如果列表中存在冒号,那么会索引列表中对应位置维度的元素。示例代码如下:

1
2
3
4
5
6
7
8
9
tensor([[4, 2, 5, 9],
[9, 1, 4, 9],
[1, 1, 7, 2]])

# 第0维中的所有元素,第1维中下标为0和1的元素
print(t2[:, [0,1]])
tensor([[4, 2],
[9, 1],
[1, 1]])

二维张量中使用多列表索引,那么获取的是配对索引。示例代码如下:

1
2
3
4
5
6
tensor([[4, 2, 5, 9],
[9, 1, 4, 9],
[1, 1, 7, 2]])
# 获取第0行第1列的元素、获取第2行第1列的元素
print(t2[[0,2], [1,3]])
tensor([2, 2])

以上是获取(0,1)和(2,3)两个位置的元素,不是获取第0行和第2行,第1列和第3列的元素,如果想要获取后者,那么代码应该修改为:

1
2
3
4
# 先把行取出来,再取列  
print(t1[[0,1]][:, [2,1]])
tensor([[5, 2],
[4, 1]])

三、布尔索引

布尔索引示例代码如下:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
t1 = torch.tensor([[1, 2, 3],  
                  [4, 5, 6],
                  [7, 8, 9]])

# 获取张量中值大于3的元素,返回一个一维张量
mask = t1 > 3
print(mask)
tensor([[False, False, False],
[ True, True, True],
[ True, True, True]])

print(t1[mask])
tensor([4, 5, 6, 7, 8, 9])

# 针对某个维度使用布尔索引
# 获取第1个维度大于4的元素
mask = t1[:, 1] > 4
print(mask)
tensor([False, True, True])

print(t1[mask])
tensor([[4, 5, 6],
[7, 8, 9]])

# 多个条件
# 与
mask = (t1 > 3) & (t1 < 8)
# 或
mask = (t1 > 8) | (t1 < 3)
# 非
mask = ~(t1 == 5) #也可以写成 t1 != 5


作者

步步为营

发布于

2026-09-29

更新于

2026-09-29

许可协议