Skip to content

NumPy - 高级索引

除了基本切片(slicing)(例如,arr[start:stop:step])之外,NumPy 还提供了更复杂的方式来选择数组中的元素,统称为 高级索引(Advanced Indexing)。这通常涉及使用索引数组(整数数组)或布尔数组(mask,掩码)来访问数组元素。

与基本切片的一个关键区别在于,高级索引总是返回数据的 副本(copy),而不是视图(view)。这意味着对通过高级索引创建的新数组的更改 不会 影响原始数组。

高级索引主要有两种类型:

  • 整数数组索引(Integer Array Indexing)
  • 布尔数组索引(Boolean Array Indexing)

整数数组索引(Integer Array Indexing)

Section titled “整数数组索引(Integer Array Indexing)”

这允许根据使用整数数组指定的 N 维坐标来选择数组中的任意项。

如果提供的索引数组数量与目标数组的维度数量相同,NumPy 会将每个索引数组中对应的元素配对,形成所选元素的坐标。

import numpy as np
x = np.array([[1, 2, 3],
[4, 5, 6],
[7, 8, 9]])
# Define arrays for row and column indices
row_indices = np.array([0, 1, 2])
col_indices = np.array([0, 1, 0]) # Select elements at (0,0), (1,1), (2,0)
y = x[row_indices, col_indices]
print(f"Original array (x):\n{x}\n")
print(f"Selected elements (y): {y}")
# Modify the new array 'y'. This does NOT change 'x' because 'y' is a copy.
y[0] = 99
print(f"Modified selected elements (y): {y}")
print(f"Original array (x) remains unchanged:\n{x}")

输出:

Original array (x):
[[1 2 3]
[4 5 6]
[7 8 9]]
Selected elements (y): [1 5 7]
Modified selected elements (y): [99 5 7]
Original array (x) remains unchanged:
[[1 2 3]
[4 5 6]
[7 8 9]]

结果数组的形状与广播后的索引数组的形状一致。

示例 2:选择 2D 数组的角部元素

Section titled “示例 2:选择 2D 数组的角部元素”
import numpy as np
x = np.arange(12).reshape(4, 3)
print(f"Original array (x):\n{x}\n")
# Indices for top-left, top-right, bottom-left, bottom-right corners
# (0,0), (0,2), (3,0), (3,2)
row_indices = np.array([[0, 0], [3, 3]]) # Shape (2, 2)
col_indices = np.array([[0, 2], [0, 2]]) # Shape (2, 2)
corners = x[row_indices, col_indices]
print(f"Corner elements (shape matches index arrays):\n{corners}")

输出:

Original array (x):
[[ 0 1 2]
[ 3 4 5]
[ 6 7 8]
[ 9 10 11]]
Corner elements (shape matches index arrays):
[[ 0 2]
[ 9 11]]

可以将基本切片与整数数组索引混合使用。切片选择范围,而整数数组选择这些范围内的特定索引(或其他维度)上的元素。

import numpy as np
x = np.arange(12).reshape(4, 3)
print(f"Original array (x):\n{x}\n")
# Select rows 1 and 3, and column 2
y = x[[1, 3], 2]
print(f"Rows 1 & 3, column 2: {y}")
# Select row 2, and columns 0 and 2
z = x[2, [0, 2]]
print(f"Row 2, columns 0 & 2: {z}")
# Select all rows (:), and columns 0 and 2
w = x[:, [0, 2]]
print(f"All rows, columns 0 & 2:\n{w}")

输出:

Original array (x):
[[ 0 1 2]
[ 3 4 5]
[ 6 7 8]
[ 9 10 11]]
Rows 1 & 3, column 2: [ 5 11]
Row 2, columns 0 & 2: [6 8]
All rows, columns 0 & 2:
[[ 0 2]
[ 3 5]
[ 6 8]
[ 9 11]]

布尔数组索引(Boolean Array Indexing)

Section titled “布尔数组索引(Boolean Array Indexing)”

这涉及使用与原始数组形状相同的布尔数组(一个 mask,掩码)。对应于掩码中 True 值的元素会被选中,而对应于 False 值的元素会被丢弃。结果通常是一个只包含所选元素的 1 维数组。

import numpy as np
x = np.arange(12).reshape(4, 3)
print(f"Original array (x):\n{x}\n")
# Create a boolean mask for elements greater than 5
mask = x > 5
print(f"Boolean mask (x > 5):\n{mask}\n")
# Apply the mask
selected = x[mask]
print(f"Elements greater than 5: {selected}")
# You can apply the condition directly
print(f"Elements less than or equal to 5: {x[x <= 5]}")

输出:

Original array (x):
[[ 0 1 2]
[ 3 4 5]
[ 6 7 8]
[ 9 10 11]]
Boolean mask (x > 5):
[[False False False]
[False False False]
[ True True True]
[ True True True]]
Elements greater than 5: [ 6 7 8 9 10 11]
Elements less than or equal to 5: [0 1 2 3 4 5]

布尔索引对于过滤掉 NaN (Not a Number,非数字) 值很有用。

import numpy as np
a = np.array([np.nan, 1., 2., np.nan, 3., 4., 5.])
print(f"Array with NaNs: {a}")
# Create mask for non-NaN values using np.isnan and the ~ (NOT) operator
mask = ~np.isnan(a)
print(f"Is not NaN mask: {mask}")
# Select non-NaN elements
filtered_a = a[mask]
print(f"Array without NaNs: {filtered_a}")
# Direct filtering
print(f"Directly filtered: {a[~np.isnan(a)]}")

输出:

Array with NaNs: [nan 1. 2. nan 3. 4. 5.]
Is not NaN mask: [False True True False True True True]
Array without NaNs: [1. 2. 3. 4. 5.]
Directly filtered: [1. 2. 3. 4. 5.]

可以使用 np.iscomplex 等函数创建用于过滤特定数据类型的掩码。

import numpy as np
a = np.array([1, 2 + 6j, 5, 3.5 + 5j, 0j])
print(f"Mixed type array: {a}")
# Select only complex numbers
complex_mask = np.iscomplex(a)
print(f"Is complex mask: {complex_mask}")
print(f"Complex elements only: {a[complex_mask]}")

输出:

Mixed type array: [1. +0.j 2. +6.j 5. +0.j 3.5+5.j 0. +0.j]
Is complex mask: [False True False True False]
Complex elements only: [2. +6.j 3.5+5.j]

应用:高级索引对于数据过滤、条件选择和修改、重排元素以及实现需要访问数组特定子集的复杂算法至关重要。

有关全面概述,请查阅 NumPy 索引文档。