Administrator
发布于 2026-09-09 / 0 阅读
0
0

Python TP3:NumPy 的 dtype、strides、视图与副本

NumPy中的dtype、strides、视图和副本

视图还是副本(基础)

问题背景
x = np.arange(12, dtype=np.float64)
A = x.reshape(3, 4)
B = A[:, 0]

这段代码创建了一个包含0到11的一维数组 x,数据类型为64位浮点数(每个元素占8字节)。然后将 x 重塑为3行4列的二维数组 A,最后提取 A 的第一列赋给 B

B.base is A 的结果
print(f"B.base is A: {B.base is A}")  # False
print(f"B.base is x: {B.base is x}")  # True

执行结果:

B.base is A: False
B.base is x: True

B.base 返回的是 B 所依赖的最原始的数组。由于 A 本身也是 x 的视图(reshape 在可能的情况下返回视图),B 的真正基底是最底层的 x,而不是中间层的 A.base 属性总是指向内存的实际拥有者。

B的strides分析
print(f"x.shape: {x.shape}, x.strides: {x.strides}")  # (12,), (8,)
print(f"A.shape: {A.shape}, A.strides: {A.strides}")  # (3, 4), (32, 8)
print(f"B.shape: {B.shape}, B.strides: {B.strides}")  # (3,), (32,)

执行结果:

x.shape: (12,), x.strides: (8,)
A.shape: (3, 4), A.strides: (32, 8)
B.shape: (3,), B.strides: (32,)

strides表示在内存中从一个元素移动到下一个元素需要跨越的字节数。float64 类型每个元素占8字节。

对于一维数组 x,相邻元素间隔8字节,因此 strides = (8,)

对于二维数组 A(3行4列),从一行移动到下一行需要跨越4个元素即 4 \times 8 = 32 字节,同一行内相邻元素间隔8字节,因此 strides = (32, 8)

对于 B = A[:, 0](提取第一列),B 包含 A[0,0]A[1,0]A[2,0] 三个元素。这些元素在原始内存中每隔一行出现,间隔32字节,因此 strides = (32,)。这个非连续的stride正是视图能够工作的关键。

修改B对A的影响

执行 B[:] = -1 后会发生什么?

print(f"修改前 A:\n{A}")
B[:] = -1
print(f"修改后 A:\n{A}")
print(f"修改后 x: {x}")

执行结果:

修改前 A:
[[ 0.  1.  2.  3.]
 [ 4.  5.  6.  7.]
 [ 8.  9. 10. 11.]]
修改后 A:
[[-1.  1.  2.  3.]
 [-1.  5.  6.  7.]
 [-1.  9. 10. 11.]]
修改后 x: [-1.  1.  2.  3. -1.  5.  6.  7. -1.  9. 10. 11.]

B 是视图而非副本,它与 Ax 共享同一块内存。修改 B 的元素会同时修改 A 的第一列以及 x 中对应位置(索引0、4、8)的元素。

内存共享的证明
x = np.arange(12, dtype=np.float64)
A = x.reshape(3, 4)
B = A[:, 0]

print(f"np.shares_memory(B, A): {np.shares_memory(B, A)}")  # True
print(f"np.shares_memory(B, x): {np.shares_memory(B, x)}")  # True
print(f"B 的内存地址: {B.__array_interface__['data']}")
print(f"A 的内存地址: {A.__array_interface__['data']}")
print(f"x 的内存地址: {x.__array_interface__['data']}")

执行结果:

np.shares_memory(B, A): True
np.shares_memory(B, x): True
B 的内存地址: 218201728
A 的内存地址: 218201728
x 的内存地址: 218201728

np.shares_memory() 函数检查两个数组是否共享内存。三个数组的起始内存地址相同,证明它们指向同一块内存区域。__array_interface__ 是NumPy数组的底层接口,其中 'data' 键包含内存地址信息。

高级索引

问题背景
x = np.arange(10)
idx = [0, 2, 4, 6]
y = x[idx]

这里使用列表 idx 作为索引来提取 x 中的特定元素。这种使用列表或数组作为索引的方式称为高级索引(advanced indexing)或花式索引(fancy indexing)。

y是视图还是副本
print(f"np.shares_memory(x, y): {np.shares_memory(x, y)}")  # False

执行结果:

np.shares_memory(x, y): False

y 是副本而非视图。使用列表或数组进行索引(高级索引)总是返回副本,因为所选元素在内存中的位置不一定是等间距的,无法用简单的strides来表示。

y.base的值
print(f"y.base: {y.base}")  # None

执行结果:

y.base: None

由于 y 是副本,它拥有自己独立的内存,不依赖于任何其他数组,因此 base 属性为 None。这是判断数组是否为副本的一种方法:如果 arr.base is None 且数组不是标量,则该数组拥有自己的数据。

修改y的效果

y[:] = 100 的效果是什么?

print(f"修改前 x: {x}")
print(f"修改前 y: {y}")
y[:] = 100
print(f"修改后 y: {y}")
print(f"修改后 x: {x}")

执行结果:

修改前 x: [0 1 2 3 4 5 6 7 8 9]
修改前 y: [0 2 4 6]
修改后 y: [100 100 100 100]
修改后 x: [0 1 2 3 4 5 6 7 8 9]

由于 y 是副本,修改 y 不会影响原始数组 x。两者的内存完全独立。

与切片索引的区别
x = np.arange(10)
z = x[::2]  # 切片索引

print(f"z = x[::2]: {z}")
print(f"np.shares_memory(x, z): {np.shares_memory(x, z)}")  # True
print(f"z.base is x: {z.base is x}")  # True

z[:] = -1
print(f"修改 z 后 x: {x}")

执行结果:

z = x[::2]: [0 2 4 6]
np.shares_memory(x, z): True
z.base is x: True
修改 z 后 x: [-1  1 -1  3 -1  5 -1  7 -1  9]

x[::2] 使用切片索引,返回视图;x[idx] 使用列表索引,返回副本。虽然两者提取的元素相同(索引0、2、4、6),但行为完全不同。

根本原因在于内存布局的可表示性。切片 [::2] 选取的元素在内存中是等间距分布的(每隔2个元素取1个),可以用strides来描述这种访问模式。z 的strides为 (16,)(假设 int64,每个元素8字节,间隔2个元素即16字节),NumPy可以创建视图。

而列表索引 [0, 2, 4, 6] 虽然在这个例子中恰好是等间距的,但NumPy无法在一般情况下假设列表索引是等间距的(列表可以是任意的如 [0, 3, 7, 9])。因此,使用列表索引时,NumPy必须复制数据来创建一个新的连续数组。

这是NumPy设计中视图与副本的核心区别:能用strides表示的访问模式返回视图,否则返回副本。

非连续切片

问题背景
x = np.arange(20, dtype=np.int64)
y = x[::2]

创建一个包含0到19的一维数组 x,数据类型为64位整数(每个元素占8字节)。然后使用步长为2的切片 [::2] 提取偶数索引位置的元素。

y的strides分析
print(f"x.strides: {x.strides}")  # (8,)
print(f"y.strides: {y.strides}")  # (16,)
print(f"y.shape: {y.shape}")      # (10,)

执行结果:

x.strides: (8,)
y.strides: (16,)
y.shape: (10,)

x 中相邻元素间隔8字节(int64 的大小)。y = x[::2] 每隔一个元素取一个,因此 y 中相邻元素在原始内存中实际间隔2个元素,即 2 \times 8 = 16 字节。y 包含10个元素(索引0、2、4、...、18)。

y的内存连续性
print(f"y.flags['C_CONTIGUOUS']: {y.flags['C_CONTIGUOUS']}")  # False
print(f"x.flags['C_CONTIGUOUS']: {x.flags['C_CONTIGUOUS']}")  # True

执行结果:

y.flags['C_CONTIGUOUS']: False
x.flags['C_CONTIGUOUS']: True

C_CONTIGUOUS 标志表示数组元素在内存中是否按C语言风格(行优先)连续存储。x 是连续的,元素紧密排列。y 不是连续的,因为它的元素在内存中被跳过的元素间隔开(0、2、4...之间夹着1、3、5...)。

NumPy避免复制数据的原因
print(f"np.shares_memory(x, y): {np.shares_memory(x, y)}")  # True
print(f"y.base is x: {y.base is x}")  # True

执行结果:

np.shares_memory(x, y): True
y.base is x: True

NumPy尽可能避免复制数据,原因有三:复制需要额外的内存分配和数据传输时间,对于大数组开销显著;strides机制可以在不复制的情况下表示各种访问模式,包括非连续访问;视图允许对原始数据进行原地修改,这在很多算法中是必需的。这是NumPy高效的核心设计原则。

修改y对x的影响
x = np.arange(20, dtype=np.int64)
y = x[::2]
print(f"修改前 x: {x}")
y[:] = 0
print(f"修改后 x: {x}")
print(f"修改后 y: {y}")

执行结果:

修改前 x: [ 0  1  2  3  4  5  6  7  8  9 10 11 12 13 14 15 16 17 18 19]
修改后 x: [ 0  1  0  3  0  5  0  7  0  9  0 11  0 13  0 15  0 17  0 19]
修改后 y: [0 0 0 0 0 0 0 0 0 0]

y 是视图,修改 y 会将 x 中偶数索引位置(0、2、4、...、18)的元素都变为0,而奇数索引位置的元素保持不变。

转置

问题背景
A = np.arange(12).reshape(3, 4)
B = A.T

创建一个3行4列的数组 A,然后通过 .T 属性获取其转置 B。转置将行变成列、列变成行,得到4行3列的数组。

B是否为视图
print(f"np.shares_memory(A, B): {np.shares_memory(A, B)}")  # True
print(f"B.base is A.base: {B.base is A.base}")  # True

执行结果:

np.shares_memory(A, B): True
B.base is A.base: True

B 是视图而非副本。转置操作不复制任何数据,它只是改变了数组的shape和strides来重新解释同一块内存。这使得转置成为一个 O(1) 的常数时间操作,无论数组多大。

A和B的strides比较
print(f"A.shape: {A.shape}")        # (3, 4)
print(f"A.strides: {A.strides}")    # (32, 8)
print(f"B.shape: {B.shape}")        # (4, 3)
print(f"B.strides: {B.strides}")    # (8, 32)

执行结果:

A.shape: (3, 4)
A.strides: (32, 8)
B.shape: (4, 3)
B.strides: (8, 32)

A 的strides为 (32, 8):沿第一个轴(行方向)移动需要跨越32字节(一行4个元素,每个8字节),沿第二个轴(列方向)移动需要跨越8字节。这是C顺序(行优先)的典型特征:行内元素连续存储。

B 的strides为 (8, 32):strides被交换了。在 B 中,沿第一个轴移动只需8字节,沿第二个轴移动需要32字节。转置仅仅是交换了shape和strides的顺序,内存中的数据完全没有移动。

B的连续性
print(f"A.flags['C_CONTIGUOUS']: {A.flags['C_CONTIGUOUS']}")  # True
print(f"A.flags['F_CONTIGUOUS']: {A.flags['F_CONTIGUOUS']}")  # False
print(f"B.flags['C_CONTIGUOUS']: {B.flags['C_CONTIGUOUS']}")  # False
print(f"B.flags['F_CONTIGUOUS']: {B.flags['F_CONTIGUOUS']}")  # True

执行结果:

A.flags['C_CONTIGUOUS']: True
A.flags['F_CONTIGUOUS']: False
B.flags['C_CONTIGUOUS']: False
B.flags['F_CONTIGUOUS']: True

A 是C连续的(行优先,C语言风格),行内元素在内存中连续存储。B 不是C连续的,但是F连续的(列优先,Fortran风格)。F连续意味着列内元素在内存中连续存储。

这是因为 A 按行存储时,每一行是连续的。转置后,原来的行变成了列,所以 B 的列是连续的,即F连续。这两种连续性在不同的计算库中有性能影响:C语言和Python偏好C顺序,Fortran和MATLAB偏好F顺序。

修改B对A的影响

B[0, 0] = -99 后会发生什么?

A = np.arange(12).reshape(3, 4)
B = A.T
print(f"修改前 A:\n{A}")
print(f"修改前 B:\n{B}")
B[0, 0] = -99
print(f"修改后 B:\n{B}")
print(f"修改后 A:\n{A}")

执行结果:

修改前 A:
[[ 0  1  2  3]
 [ 4  5  6  7]
 [ 8  9 10 11]]
修改前 B:
[[ 0  4  8]
 [ 1  5  9]
 [ 2  6 10]
 [ 3  7 11]]
修改后 B:
[[-99   4   8]
 [  1   5   9]
 [  2   6  10]
 [  3   7  11]]
修改后 A:
[[-99   1   2   3]
 [  4   5   6   7]
 [  8   9  10  11]]

B 是视图,B[0, 0]A[0, 0] 指向同一块内存。修改 B[0, 0] 会同时改变 A[0, 0] 的值。在转置关系中,B[i, j] 总是与 A[j, i] 共享同一个内存位置。

静默reshape

问题背景
x = np.arange(12)
y = x[::2]
z = y.reshape(2, 3)

创建一个包含0到11的数组 x,使用步长为2的切片得到 y(包含6个元素:0、2、4、6、8、10),然后尝试将 y 重塑为2行3列的数组。

代码有效性
print(f"x: {x}")
print(f"y = x[::2]: {y}")
print(f"z = y.reshape(2, 3):\n{z}")

执行结果:

x: [ 0  1  2  3  4  5  6  7  8  9 10 11]
y = x[::2]: [ 0  2  4  6  8 10]
z = y.reshape(2, 3):
[[ 0  2  4]
 [ 6  8 10]]

代码有效。y 有6个元素,6 = 2 \times 3,可以reshape成 (2, 3) 的形状。

z是视图还是副本
print(f"np.shares_memory(y, z): {np.shares_memory(y, z)}")  # False
print(f"np.shares_memory(x, z): {np.shares_memory(x, z)}")  # False
print(f"z.base: {z.base}")  # None

执行结果:

np.shares_memory(y, z): False
np.shares_memory(x, z): False
z.base: None

z 是副本,不与 yx 共享内存。z.baseNone 表明 z 拥有自己独立的内存。这种情况下 reshape 静默地创建了副本,没有任何警告或错误。

NumPy被迫复制的原因
print(f"y.strides: {y.strides}")  # (16,)
print(f"y.flags['C_CONTIGUOUS']: {y.flags['C_CONTIGUOUS']}")  # False

执行结果:

y.strides: (16,)
y.flags['C_CONTIGUOUS']: False

y 是非连续的,stride为16字节(每隔一个 int64 元素)。reshape成 (2, 3) 要求数据在内存中连续排列,即第一行的3个元素紧密相邻,第二行的3个元素也紧密相邻。

然而 y 的元素在内存中是分散的(位置0、2、4、6、8、10),无法用任何strides组合来表示一个连续的 (2, 3) 数组。对于连续的 (2, 3) 数组,strides应该是 (24, 8)(每行3个元素×8字节=24字节),但这与 y 的实际内存布局不兼容。因此NumPy被迫复制数据以满足reshape的要求。

验证副本的方法
x = np.arange(12)
y = x[::2]
z = y.reshape(2, 3)

# 方法 1: 使用 np.shares_memory
print(f"方法1 - np.shares_memory(y, z): {np.shares_memory(y, z)}")  # False

# 方法 2: 检查 base
print(f"方法2 - z.base is y: {z.base is y}")  # False
print(f"方法2 - z.base is None: {z.base is None}")  # True

# 方法 3: 修改 z 观察 y 是否改变
z[0, 0] = -999
print(f"方法3 - 修改 z 后 y: {y}")  # y 不变,证明是副本

# 方法 4: 比较内存地址
x = np.arange(12)
y = x[::2]
z = y.reshape(2, 3)
print(f"方法4 - y 的地址: {y.__array_interface__['data']}")
print(f"方法4 - z 的地址: {z.__array_interface__['data']}")

执行结果:

方法1 - np.shares_memory(y, z): False
方法2 - z.base is y: False
方法2 - z.base is None: True
方法3 - 修改 z 后 y: [ 0  2  4  6  8 10]

方法3最直观:修改 z[0, 0]y 保持不变,证明两者内存独立。方法4通过比较内存地址,地址不同也证明了副本的存在。

astype类型转换

问题背景
x = np.arange(5, dtype=np.float64)
y = x.astype(np.int32)

创建一个 float64 类型的数组 x,然后使用 astype 转换为 int32 类型的数组 y

y与x是否共享内存
print(f"np.shares_memory(x, y): {np.shares_memory(x, y)}")  # False
print(f"y.base: {y.base}")  # None

执行结果:

np.shares_memory(x, y): False
y.base: None

y 不与 x 共享内存,astype 在类型不同时创建了副本。

astype能否返回视图
# 当 dtype 相同时,astype 可以返回视图
x = np.arange(5, dtype=np.float64)
y_same = x.astype(np.float64, copy=False)

print(f"dtype 不同时 - np.shares_memory(x, x.astype(np.int32)): {np.shares_memory(x, x.astype(np.int32))}")
print(f"dtype 相同时 - np.shares_memory(x, y_same): {np.shares_memory(x, y_same)}")

执行结果:

dtype 不同时 - np.shares_memory(x, x.astype(np.int32)): False
dtype 相同时 - np.shares_memory(x, y_same): True

当目标dtype与原始dtype不同时,astype 必须创建副本,因为需要进行实际的数据转换。当目标dtype相同且指定 copy=False 时,astype 可以返回视图,避免不必要的复制。

从内存布局角度的证明
x = np.arange(5, dtype=np.float64)
y = x.astype(np.int32)

print(f"x.dtype: {x.dtype}, x.itemsize: {x.itemsize} 字节")
print(f"y.dtype: {y.dtype}, y.itemsize: {y.itemsize} 字节")

print(f"x.nbytes: {x.nbytes}")  # 总字节数
print(f"y.nbytes: {y.nbytes}")

print(f"x.strides: {x.strides}")
print(f"y.strides: {y.strides}")

执行结果:

x.dtype: float64, x.itemsize: 8 字节
y.dtype: int32, y.itemsize: 4 字节
x.nbytes: 40
y.nbytes: 20
x.strides: (8,)
y.strides: (4,)

float64 每个元素占8字节,int32 每个元素占4字节。x 占用40字节(5 \times 8),y 占用20字节(5 \times 4)。内存布局完全不同,无法用strides技巧共享内存。

更根本的原因是数据的二进制表示完全不同。float64 使用IEEE 754双精度浮点格式存储数据,而 int32 使用二进制补码整数格式。同一个数值(如3.0)在两种格式下的二进制表示完全不同,因此必须进行实际的数据转换和复制,不可能通过视图实现。

不复制地更改dtype

问题背景
x = np.array([1, 2, 3, 4], dtype=np.int32)
y = x.view(np.int16)

创建一个 int32 类型的数组 x,然后使用 view 方法以 int16 类型重新解释同一块内存。viewastype 不同,它不进行数据转换,而是直接改变对内存的解释方式。

比较nbytes
print(f"x.nbytes: {x.nbytes}")  # 16 字节 (4 * 4)
print(f"y.nbytes: {y.nbytes}")  # 16 字节 (8 * 2)

执行结果:

x.nbytes: 16
y.nbytes: 16

总字节数完全相同,都是16字节。这是因为 view 不复制数据,xy 指向同一块16字节的内存区域,只是用不同的方式来解释它。

比较形状
print(f"x.shape: {x.shape}")      # (4,)
print(f"y.shape: {y.shape}")      # (8,)
print(f"x.itemsize: {x.itemsize}")  # 4 字节
print(f"y.itemsize: {y.itemsize}")  # 2 字节

执行结果:

x.shape: (4,)
y.shape: (8,)
x.itemsize: 4
y.itemsize: 2

x 包含4个 int32 元素,每个占4字节。y 包含8个 int16 元素,每个占2字节。总字节数相同(4 \times 4 = 8 \times 2 = 16),但元素数量翻倍。NumPy自动调整了形状以保持总内存大小不变。

y实际代表什么
print(f"x: {x}")
print(f"y: {y}")

执行结果:

x: [1 2 3 4]
y: [1 0 2 0 3 0 4 0]

yx 的每个 int32(4字节)重新解释为两个 int16(各2字节)。以 x[0] = 1 为例,在小端序(little-endian)系统中,整数1的 int32 表示在内存中是 [0x01, 0x00, 0x00, 0x00](低字节在前)。这4个字节被解释为两个 int16:前两个字节 [0x01, 0x00] 是1,后两个字节 [0x00, 0x00] 是0。

y 不是 x 的数学转换,而是对同一块内存的不同解释。数值1、2、3、4在 int32 格式下高16位都是0,所以 y 中交替出现原始值和0。

内存层面的详细描述
print(f"x 的内存地址: {x.__array_interface__['data']}")
print(f"y 的内存地址: {y.__array_interface__['data']}")
print(f"np.shares_memory(x, y): {np.shares_memory(x, y)}")  # True

# 查看原始字节
print(f"x 的原始字节: {x.tobytes().hex()}")
print(f"y 的原始字节: {y.tobytes().hex()}")

print(f"x.strides: {x.strides}")  # (4,)
print(f"y.strides: {y.strides}")  # (2,)

执行结果:

np.shares_memory(x, y): True
x 的原始字节: 01000000020000000300000004000000
y 的原始字节: 01000000020000000300000004000000
x.strides: (4,)
y.strides: (2,)

xy 指向完全相同的内存地址,原始字节序列完全一致。view() 只是创建了一个新的数组对象,使用不同的dtype、shape和strides来解释同一块内存。没有任何数据被复制或移动。

strides也相应调整:x 的stride是4字节(一个 int32 的大小),y 的stride是2字节(一个 int16 的大小)。

潜在的严重bug
x = np.array([1, 2, 3, 4], dtype=np.int32)
y = x.view(np.int16)

print(f"修改前 x: {x}")
print(f"修改前 y: {y}")

# Bug 1: 修改 y 会意外修改 x
y[0] = 999
print(f"修改 y[0]=999 后 x: {x}")
print(f"修改 y[0]=999 后 y: {y}")

执行结果:

修改前 x: [1 2 3 4]
修改前 y: [1 0 2 0 3 0 4 0]
修改 y[0]=999 后 x: [999   2   3   4]
修改 y[0]=999 后 y: [999   0   2   0   3   0   4   0]

view 可能导致多种严重bug:

共享内存导致意外副作用是最常见的问题。修改 y[0] 会改变 x[0] 的低16位。如果代码的其他部分依赖 x 的值保持不变,就会产生难以追踪的bug。

数值含义被改变也是隐患。y 的值看起来像是有意义的数字(1、0、2、0...),但实际上是字节级的重新解释,不是数学上的类型转换。如果误将 y 当作正常的整数数组使用,会导致逻辑错误。

字节序依赖性影响可移植性。上述结果基于小端序系统。在大端序系统上,相同代码会产生不同结果:y 可能是 [0, 1, 0, 2, 0, 3, 0, 4]。这种平台依赖性使代码难以跨系统运行。

调试困难是这类bug的特点。当一个数组的修改莫名影响另一个看似无关的数组时,问题根源很难定位,特别是在大型代码库中。

因此,view 应该谨慎使用,主要用于底层内存操作或性能关键的场景,且必须清楚了解其内存共享的特性。

字节级解释

问题背景
x = np.arange(6, dtype=np.int32)
y = x.view(np.uint8)

创建一个包含0到5的 int32 数组 x,然后使用 viewuint8(无符号8位整数,即单字节)重新解释同一块内存。

y的大小
print(f"x: {x}")
print(f"x.shape: {x.shape}")  # (6,)
print(f"x.size: {x.size}")    # 6
print(f"y.shape: {y.shape}")  # (24,)
print(f"y.size: {y.size}")    # 24

执行结果:

x: [0 1 2 3 4 5]
x.shape: (6,)
x.size: 6
y.shape: (24,)
y.size: 24

x 有6个 int32 元素,每个占4字节,总共24字节。y 将同一块24字节的内存解释为 uint8,每个元素占1字节,因此 y 有24个元素。元素数量的关系为 y.\text{size} = x.\text{size} \times 4 = 6 \times 4 = 24

与itemsize的联系
print(f"x.itemsize: {x.itemsize}")  # 4 字节
print(f"y.itemsize: {y.itemsize}")  # 1 字节
print(f"x.nbytes: {x.nbytes}")      # 24 字节
print(f"y.nbytes: {y.nbytes}")      # 24 字节

print(f"y.size = x.size * (x.itemsize / y.itemsize)")
print(f"y.size = {x.size} * ({x.itemsize} / {y.itemsize}) = {x.size * (x.itemsize // y.itemsize)}")

执行结果:

x.itemsize: 4
y.itemsize: 1
x.nbytes: 24
y.nbytes: 24
y.size = x.size * (x.itemsize / y.itemsize)
y.size = 6 * (4 / 1) = 24

itemsize 是每个元素占用的字节数。int32itemsize 为4字节,uint8itemsize 为1字节。总字节数(nbytes)不变,都是24字节。元素数量由总字节数除以单个元素大小决定:

y.\text{size} = \frac{x.\text{nbytes}}{y.\text{itemsize}} = \frac{24}{1} = 24
y[0]对应什么
print(f"x: {x}")
print(f"y: {y}")
print(f"x[0] = {x[0]}")
print(f"y[0:4] = {y[0:4]}")  # x[0] 的 4 个字节

执行结果:

x: [0 1 2 3 4 5]
y: [0 0 0 0 1 0 0 0 2 0 0 0 3 0 0 0 4 0 0 0 5 0 0 0]
x[0] = 0
y[0:4] = [0 0 0 0]

y[0] 对应 x[0] 的最低有效字节。在小端序(little-endian)系统中,低位字节存储在低地址。

x[0] = 0 在内存中表示为 [0x00, 0x00, 0x00, 0x00],对应 y[0:4] = [0, 0, 0, 0]

x[1] = 1 在内存中表示为 [0x01, 0x00, 0x00, 0x00],对应 y[4:8] = [1, 0, 0, 0]

x[5] = 5 在内存中表示为 [0x05, 0x00, 0x00, 0x00],对应 y[20:24] = [5, 0, 0, 0]

由于这些数值都很小(0到5),只有最低字节非零,其余三个字节都是0。

strides的确切作用
print(f"x.strides: {x.strides}")  # (4,)
print(f"y.strides: {y.strides}")  # (1,)

执行结果:

x.strides: (4,)
y.strides: (1,)

strides定义了在内存中从一个元素移动到下一个元素需要跳过的字节数。

对于 xint32),strides = (4,) 表示从 x[i]x[i+1] 需要跳过4字节,因为每个 int32 占4字节。

对于 yuint8),strides = (1,) 表示从 y[i]y[i+1] 需要跳过1字节,因为每个 uint8 占1字节。

内存布局示意:

地址:    0    1    2    3    4    5    6    7    ...
x:      [    x[0]=0     ] [    x[1]=1     ]  ...
y:      y[0] y[1] y[2] y[3] y[4] y[5] y[6] y[7] ...

x[0] 占据地址0-3,x[1] 占据地址4-7。y[0]y[3] 分别占据地址0、1、2、3,恰好覆盖 x[0] 的4个字节。strides确保了无论用哪种dtype解释,都能正确遍历内存。

实际的静默bug

问题背景
A = np.ones((4, 4))
B = A[::2, ::2]
B *= 10

创建一个4×4的全1数组 A,使用步长为2的切片在两个维度上提取元素得到 B(2×2数组),然后对 B 进行原地乘法。

A是否被修改
print(f"修改前 A:\n{A}")
print(f"修改前 B:\n{B}")
print(f"np.shares_memory(A, B): {np.shares_memory(A, B)}")  # True

B *= 10

print(f"\n执行 B *= 10 后:")
print(f"修改后 B:\n{B}")
print(f"修改后 A:\n{A}")

执行结果:

修改前 A:
[[1. 1. 1. 1.]
 [1. 1. 1. 1.]
 [1. 1. 1. 1.]
 [1. 1. 1. 1.]]
修改前 B:
[[1. 1.]
 [1. 1.]]
np.shares_memory(A, B): True

执行 B *= 10 后:
修改后 B:
[[10. 10.]
 [10. 10.]]
修改后 A:
[[10.  1. 10.  1.]
 [ 1.  1.  1.  1.]
 [10.  1. 10.  1.]
 [ 1.  1.  1.  1.]]

A 被修改了。BA 的视图(切片索引返回视图),两者共享内存。B *= 10 是原地操作,修改了 B 指向的内存,而这块内存同时也是 A 的一部分。A 的位置 [0,0][0,2][2,0][2,2] 被改成了10。

这种行为的危险性
print(f"B.base is A: {B.base is A}")  # True
print(f"B.flags['OWNDATA']: {B.flags['OWNDATA']}")  # False

执行结果:

B.base is A: True
B.flags['OWNDATA']: False

这种行为危险的原因有多个方面。静默修改是最主要的问题:代码看起来只修改 B,但实际上 A 也被修改了,没有任何警告或错误提示。难以追踪也是问题所在:在复杂代码中,很难意识到 BA 的视图,特别是当变量名和操作分布在不同函数或文件中时。意外的副作用会导致bug:函数接收数组参数时可能意外修改调用者的数据。非显式的语法也容易误导:B = A[::2, ::2] 看起来像是创建新数组,实际上是视图。调试困难是最终结果:bug可能在远离实际原因的地方显现。

使用 .copy() 可以避免这种问题:

A = np.ones((4, 4))
B = A[::2, ::2].copy()
B *= 10
print(f"B:\n{B}")
print(f"A:\n{A}")  # A 保持不变
实际应用场景

场景一是图像处理中的下采样。创建缩略图进行预览处理时,如果不小心使用了视图,对缩略图的操作会破坏原始图像:

image = np.random.randint(0, 256, (100, 100), dtype=np.uint8)
original_sum = image.sum()

thumbnail = image[::10, ::10]  # 视图!
thumbnail[:] = 0  # "清空"缩略图进行测试

print(f"原始图像总和变化: {original_sum} -> {image.sum()}")
# 原始图像被意外破坏了!

场景二是科学计算中的数据归一化。对数据子集进行归一化时,可能意外归一化了原始数据:

raw_data = np.array([[100, 200], [300, 400]], dtype=np.float64)
subset = raw_data[::1, ::1]  # 看起来像副本,实际是视图
subset /= subset.max()  # 归一化"副本"
# 原始数据被意外归一化了!

场景三是机器学习中的训练/测试集划分。中心化训练集时可能污染整个数据集:

dataset = np.arange(20).reshape(4, 5).astype(float)
train_set = dataset[::2]  # 视图
train_set -= train_set.mean()  # 中心化训练集
# 原始数据集被污染了!测试集也受影响!

场景四是金融数据处理。对采样数据进行调整时可能篡改历史记录:

stock_prices = np.array([100.0, 102.0, 98.0, 105.0, 110.0])
weekly_sample = stock_prices[::2]  # 周采样
weekly_sample *= 1.1  # 模拟 10% 调整
# 历史数据被篡改了!

这些场景的共同点是:开发者以为在操作数据的副本,实际上操作的是视图,导致原始数据被意外修改。解决方案是在不确定时显式调用 .copy(),或使用 np.shares_memory() 检查内存共享关系。

诊断

对于任何NumPy数组,可以通过以下代码检查其内存模型信息:

print("shape    :", x.shape)
print("dtype    :", x.dtype)
print("strides  :", x.strides)
print("base     :", x.base is not None)
print("C contig:", x.flags['C_CONTIGUOUS'])
诊断函数实现
def diagnose(arr, name="x"):
    """诊断函数:显示数组的内存模型信息"""
    print(f"=== 诊断 {name} ===")
    print(f"shape    : {arr.shape}")
    print(f"dtype    : {arr.dtype}")
    print(f"strides  : {arr.strides}")
    print(f"base     : {arr.base is not None}")
    print(f"C contig : {arr.flags['C_CONTIGUOUS']}")
    print()
各属性含义详解

B = A[:, 0] 为例进行逐行解释:

shape 表示数组的维度和每个维度的大小。(3,) 表示一维数组,有3个元素。对于二维数组 (3, 4) 表示3行4列。shape决定了数组的逻辑结构。

dtype 表示数据类型,决定每个元素的内存表示和占用空间。float64 表示64位浮点数,每个元素占8字节。常见的dtype还有 int32(4字节整数)、uint8(1字节无符号整数)等。dtype决定了内存占用和数值精度。

strides 表示在每个维度上移动一个元素需要跳过的字节数。(32,) 表示从 B[i]B[i+1] 需要跳过32字节。这是因为 B 取的是 A 的第一列,A 每行有4个 float64 元素,共 4 \times 8 = 32 字节。strides决定了内存访问模式。

base is not None 表示该数组是否是另一个数组的视图。True 表示是视图,有一个base数组,与之共享内存;False 表示拥有自己的数据(副本或原始数组)。这是判断内存共享的关键。

C_CONTIGUOUS 表示数据在内存中是否按C顺序(行优先)连续存储。False 表示数据不连续,元素之间有间隔。这影响缓存效率和某些操作的性能,很多底层库(如BLAS)要求数据连续。

不同数组的对比分析
# 1. 原始数组
x = np.arange(12, dtype=np.float64)
# shape=(12,), strides=(8,), base=False, C_contig=True
# 解释: 连续的一维数组,拥有自己的数据

# 2. reshape 视图
A = x.reshape(3, 4)
# shape=(3,4), strides=(32,8), base=True, C_contig=True
# 解释: 二维视图,strides=(32,8) 表示行跨32字节,列跨8字节

# 3. 列切片
B = A[:, 0]
# shape=(3,), strides=(32,), base=True, C_contig=False
# 解释: 非连续视图,stride=32 因为取每行第一个元素

# 4. 步长切片
y = x[::2]
# shape=(6,), strides=(16,), base=True, C_contig=False
# 解释: 非连续视图,stride=16 因为每隔一个元素取值

# 5. 转置
T = A.T
# shape=(4,3), strides=(8,32), base=True, C_contig=False
# 解释: strides 交换,变成 F 连续(列优先)而非 C 连续

# 6. 副本
C = A.copy()
# shape=(3,4), strides=(32,8), base=False, C_contig=True
# 解释: 独立副本,拥有自己的连续内存

# 7. 花式索引
F = x[[0, 2, 4]]
# shape=(3,), strides=(8,), base=False, C_contig=True
# 解释: 副本(花式索引总是返回副本),连续存储
内存模型总结
属性 含义 重要性
shape 数组的维度和大小 决定数组的逻辑结构
dtype 数据类型,每个元素的字节数 决定内存占用和精度
strides 每个维度移动一个元素的字节跨度 决定内存访问模式
base 是否是视图(指向另一个数组的数据) 决定是否共享内存
C_CONTIGUOUS 数据是否按行优先顺序连续存储 影响性能和兼容性

关键规则总结:切片索引(x[::2]A[:, 0])返回视图,共享内存;花式索引(x[[0,1,2]])返回副本,独立内存;reshape 和转置在可能时返回视图;astype 在类型不同时返回副本;view() 返回视图,重新解释同一块内存;copy() 显式创建副本。

NumPy小练习

获取数组的内存大小
x = np.random.randn(100, 100)

# 方法 1: 使用 nbytes 属性
print(f"方法 1 - x.nbytes: {x.nbytes} 字节")

# 方法 2: 使用 size * itemsize
print(f"方法 2 - x.size * x.itemsize: {x.size * x.itemsize} 字节")

# 方法 3: 使用 sys.getsizeof(包含对象开销)
import sys
print(f"方法 3 - sys.getsizeof(x): {sys.getsizeof(x)} 字节")

nbytes 返回数组数据占用的总字节数,等于元素数量乘以每个元素的字节大小。sys.getsizeof 返回的值略大,因为包含了Python对象的额外开销。

反转向量
x = np.array([1, 2, 3, 4, 5])
y = x[::-1]
print(f"原始: {x}")  # [1 2 3 4 5]
print(f"反转: {y}")  # [5 4 3 2 1]

[::-1] 是步长为-1的切片,从末尾向开头遍历,实现反转。返回的是视图。

找出非零元素的索引
x = np.array([1, 2, 0, 0, 4, 0])
idx = np.nonzero(x)[0]
print(f"非零元素索引: {idx}")  # [0 1 4]

np.nonzero() 返回一个元组,每个元素对应一个维度的非零索引。对于一维数组,取 [0] 获取索引数组。

创建边界为1、内部为0的数组
x = np.ones((5, 5))
x[1:-1, 1:-1] = 0
print(x)

先创建全1数组,然后用切片 [1:-1, 1:-1] 选中内部区域(排除第一行/列和最后一行/列),将其设为0。

创建对角线下方有值的矩阵
x = np.diag([1, 2, 3, 4], k=-1)
print(x)

np.diag() 创建对角矩阵,k=-1 指定在主对角线下方一行放置元素。结果是5×5矩阵,[1,0][2,1][3,2][4,3] 位置分别为1、2、3、4。

创建RGBA颜色的自定义dtype
color_dtype = np.dtype([
    ('R', np.uint8),
    ('G', np.uint8),
    ('B', np.uint8),
    ('A', np.uint8)
])

colors = np.array([(255, 0, 0, 255), (0, 255, 0, 128)], dtype=color_dtype)
print(f"第一个颜色的 R 值: {colors[0]['R']}")  # 255

结构化dtype允许为数组的每个元素定义多个命名字段,可以像字典一样按名称访问。

原地取反指定范围的元素
x = np.array([1, 4, 2, 6, 9, 5, 3, 8, 7])
x[(x >= 3) & (x <= 8)] *= -1
print(f"取反后: {x}")  # [ 1 -4  2 -6  9 -5 -3 -8 -7]

布尔索引 (x >= 3) & (x <= 8) 选中3到8之间的元素,*= -1 原地取反。

远离零取整
x = np.array([-2.5, -1.3, 0.7, 1.5, 2.9])
y = np.where(x > 0, np.ceil(x), np.floor(x))
print(f"远离零取整: {y}")  # [-3. -2.  1.  2.  3.]

正数用 ceil 向上取整,负数用 floor 向下取整,使数值远离零。

找两个数组的公共值
a = np.random.randint(0, 10, 10)
b = np.random.randint(0, 10, 10)
common = np.intersect1d(a, b)
print(f"公共值: {common}")

np.intersect1d() 返回两个数组的交集,结果是排序后的唯一值。

原地计算复合表达式
A = np.array([1.0, 2.0, 3.0, 4.0])
B = np.array([5.0, 6.0, 7.0, 8.0])

# 计算 ((A+B)*(-A/2)),不创建中间数组
np.add(A, B, out=B)       # B = A + B
np.divide(A, 2, out=A)    # A = A / 2
np.negative(A, out=A)     # A = -A
np.multiply(A, B, out=A)  # A = A * B = (-A/2) * (A+B)

使用 out 参数指定输出位置,避免创建临时数组,节省内存。

提取整数部分的四种方法
x = np.random.uniform(0, 10, 5)

# 方法 1: astype
x.astype(int)

# 方法 2: np.floor
np.floor(x)

# 方法 3: np.trunc
np.trunc(x)

# 方法 4: x - x % 1
x - x % 1

对于正数,这四种方法结果相同。对于负数,floor 向下取整而 trunc 向零取整。

使数组不可变
x = np.array([1, 2, 3, 4, 5])
x.flags.writeable = False

try:
    x[0] = 100
except ValueError as e:
    print(f"修改失败: {e}")

设置 writeable = False 后,任何修改操作都会抛出 ValueError

创建覆盖区域的结构化数组
point_dtype = np.dtype([('x', np.float64), ('y', np.float64)])
n = 5
x_coords, y_coords = np.meshgrid(np.linspace(0, 1, n), np.linspace(0, 1, n))

points = np.zeros((n, n), dtype=point_dtype)
points['x'] = x_coords
points['y'] = y_coords

meshgrid 生成网格坐标,结构化数组存储每个点的x和y坐标。

找最接近给定值的元素
x = np.array([1.2, 3.5, 5.8, 7.1, 9.4])
scalar = 6.0
idx = np.argmin(np.abs(x - scalar))
closest = x[idx]  # 5.8

计算与目标值的绝对差,用 argmin 找到最小差的索引。

使用StringIO读取文本数据
from io import StringIO
data = """1, 2, 3
4, 5, 6"""
x = np.genfromtxt(StringIO(data), delimiter=',')

StringIO 将字符串模拟为文件对象,genfromtxt 按分隔符解析数据。

随机放置元素
x = np.zeros((5, 5))
p = 8
indices = np.random.choice(x.size, p, replace=False)
np.put(x, indices, 1)

np.random.choice 生成不重复的随机索引,np.put 在这些位置放置值。

减去每行均值
x = np.random.rand(3, 4)
x_centered = x - x.mean(axis=1, keepdims=True)

keepdims=True 保持维度,使广播正确进行。结果每行均值为0。

按指定列排序
x = np.array([[3, 2, 1], [1, 4, 2], [2, 1, 3]])
n = 1
x_sorted = x[x[:, n].argsort()]

x[:, n] 提取第n列,argsort() 返回排序后的索引,用于重排整个数组的行。

继承ndarray创建命名数组
class NamedArray(np.ndarray):
    def __new__(cls, input_array, name=None):
        obj = np.asarray(input_array).view(cls)
        obj.name = name
        return obj

    def __array_finalize__(self, obj):
        if obj is None:
            return
        self.name = getattr(obj, 'name', None)

__new__ 创建实例并添加 name 属性,__array_finalize__ 确保在视图和切片操作中保留属性。

测试NamedArray类
arr = NamedArray([1, 2, 3, 4, 5], name="my_array")
print(f"数组: {arr}")
print(f"名称: {arr.name}")
print(f"类型: {type(arr)}")

执行结果:

数组: [1 2 3 4 5]
名称: my_array
类型: <class '__main__.NamedArray'>

自定义的 NamedArray 类继承了 np.ndarray 的所有功能,同时添加了 name 属性。

处理重复索引的累加

使用普通索引 x[idx] += 1 时,重复的索引只会被处理一次。要正确处理重复索引,需要使用 np.add.at

x = np.zeros(5)
idx = np.array([0, 1, 1, 2, 2, 2])  # 有重复索引

np.add.at(x, idx, 1)

print(f"结果 x: {x}")  # [1. 2. 3. 0. 0.]

执行结果:

原始 x: [0. 0. 0. 0. 0.]
索引: [0 1 1 2 2 2]
结果 x: [1. 2. 3. 0. 0.]

索引0出现1次,所以 x[0]=1;索引1出现2次,所以 x[1]=2;索引2出现3次,所以 x[2]=3np.add.at 是无缓冲的原地操作,会正确累加所有重复索引。

基于索引列表累加向量元素
X = np.array([1, 2, 3, 4, 5])
I = np.array([0, 1, 1, 2, 0])  # 索引列表
F = np.zeros(3)

np.add.at(F, I, X)

print(f"累加后 F: {F}")  # [6. 5. 4.]

执行结果:

X: [1 2 3 4 5]
I: [0 1 1 2 0]
累加后 F: [6. 5. 4.]

F[0] 累加了 X[0]=1X[4]=5,得到6;F[1] 累加了 X[1]=2X[2]=3,得到5;F[2] 累加了 X[3]=4,得到4。

对多个轴求和
x = np.random.rand(2, 3, 4, 5)
result = x.sum(axis=(-2, -1))
print(f"求和后形状: {result.shape}")  # (2, 3)

axis=(-2, -1) 指定对最后两个轴(轴2和轴3)同时求和。负数索引从末尾计数,-1是最后一个轴,-2是倒数第二个轴。

einsum等价形式

np.einsum 使用爱因斯坦求和约定,可以简洁地表达各种张量运算:

A = np.array([1, 2, 3])
B = np.array([4, 5, 6])

# 内积 (inner): 对应元素相乘再求和
np.inner(A, B)           # 32
np.einsum('i,i->', A, B) # 32

# 外积 (outer): 每对元素相乘
np.outer(A, B)           # [[4,5,6], [8,10,12], [12,15,18]]
np.einsum('i,j->ij', A, B)

# 求和 (sum)
np.sum(A)                # 6
np.einsum('i->', A)      # 6

# 逐元素乘法 (mul)
A * B                    # [4, 10, 18]
np.einsum('i,i->i', A, B)

einsum的语法:输入索引用逗号分隔,-> 后是输出索引。重复的索引表示求和(除非出现在输出中)。

获取点积的对角线
A = np.random.rand(3, 4)
B = np.random.rand(4, 3)

# 方法 1: 先计算完整矩阵乘积,再取对角线(效率低)
result1 = np.diag(np.dot(A, B))

# 方法 2: 使用einsum直接计算对角线(效率高)
result2 = np.einsum('ij,ji->i', A, B)

# 方法 3: 逐元素乘法再按行求和
result3 = (A * B.T).sum(axis=1)

方法2最高效,因为只计算需要的对角元素,避免了完整矩阵乘法。'ij,ji->i' 表示 \sum_j A_{ij} B_{ji},正是对角线元素的定义。

找出现频率最高的值
x = np.array([1, 2, 2, 3, 3, 3, 4, 4, 4, 4])
most_frequent = np.bincount(x).argmax()
print(f"出现频率最高的值: {most_frequent}")  # 4

np.bincount(x) 统计每个非负整数出现的次数,返回数组的索引i处存储值i出现的次数。argmax() 返回最大值的索引,即出现次数最多的值。

生命游戏实现
def life_step(x):
    """生命游戏的一步迭代"""
    neighbors = (
        np.roll(np.roll(x, 1, 0), 1, 1) +   # 左上
        np.roll(x, 1, 0) +                   # 上
        np.roll(np.roll(x, 1, 0), -1, 1) +  # 右上
        np.roll(x, 1, 1) +                   # 左
        np.roll(x, -1, 1) +                  # 右
        np.roll(np.roll(x, -1, 0), 1, 1) +  # 左下
        np.roll(x, -1, 0) +                  # 下
        np.roll(np.roll(x, -1, 0), -1, 1)   # 右下
    )
    return ((neighbors == 3) | ((x == 1) & (neighbors == 2))).astype(int)

np.roll 沿指定轴循环移动数组元素。通过8次roll操作将数组分别向8个方向移动,然后求和得到每个位置的邻居数量。生命游戏规则:活细胞有2或3个邻居则存活,死细胞有3个邻居则复活。

Bootstrap置信区间
X = np.random.randn(100)
N = 10000

bootstrap_means = np.array([
    np.random.choice(X, size=len(X), replace=True).mean()
    for _ in range(N)
])

ci_lower = np.percentile(bootstrap_means, 2.5)
ci_upper = np.percentile(bootstrap_means, 97.5)

Bootstrap方法通过有放回重采样来估计统计量的分布。重复N次采样,每次计算样本均值,然后用百分位数确定95%置信区间(2.5%和97.5%分位点)。

小波和DCT图像压缩

压缩原理

图像压缩可以通过在变换域进行阈值处理来实现。将图像从空间域变换到频率域(DCT)或多分辨率域(小波),高频系数通常幅度较小且对视觉质量影响有限,通过阈值处理移除这些系数可以减少数据量。

离散余弦变换(DCT)

对于大小为 N \times N 的图像块 f(x,y),二维DCT定义为:

C(u,v) = \alpha(u)\alpha(v) \sum_{x=0}^{N-1} \sum_{y=0}^{N-1} f(x,y) \cos\left(\frac{(2x+1)u\pi}{2N}\right) \cos\left(\frac{(2y+1)v\pi}{2N}\right)

归一化因子为:

\alpha(k) = \begin{cases} \sqrt{\frac{1}{N}}, & k = 0 \\ \sqrt{\frac{2}{N}}, & k \neq 0 \end{cases}

DCT将图像分解为不同频率的余弦基函数的线性组合。低频系数(左上角)代表图像的整体亮度和缓慢变化,高频系数(右下角)代表细节和边缘。

阈值处理

硬阈值处理直接将小于阈值的系数置零:

\tilde{C}(u,v) = \begin{cases} C(u,v), & |C(u,v)| \geq T \\ 0, & |C(u,v)| < T \end{cases}

软阈值处理在置零的同时对保留的系数进行收缩:

\tilde{C}(u,v) = \begin{cases} \text{sign}(C(u,v))(|C(u,v)| - T), & |C(u,v)| \geq T \\ 0, & |C(u,v)| < T \end{cases}

软阈值产生更平滑的结果,硬阈值保留更多细节但可能产生振铃效应。

逆DCT重建

重建图像通过IDCT获得:

\tilde{f}(x,y) = \sum_{u=0}^{N-1} \sum_{v=0}^{N-1} \alpha(u)\alpha(v)\tilde{C}(u,v) \cos\left(\frac{(2x+1)u\pi}{2N}\right) \cos\left(\frac{(2y+1)v\pi}{2N}\right)
实现代码
from scipy.fftpack import dct, idct
import pywt

def dct2(block):
    """2D DCT变换"""
    return dct(dct(block.T, norm='ortho').T, norm='ortho')

def idct2(block):
    """2D IDCT变换"""
    return idct(idct(block.T, norm='ortho').T, norm='ortho')

def hard_threshold(coeffs, T):
    """硬阈值处理"""
    return np.where(np.abs(coeffs) >= T, coeffs, 0)

def soft_threshold(coeffs, T):
    """软阈值处理"""
    return np.sign(coeffs) * np.maximum(np.abs(coeffs) - T, 0)

二维DCT通过对行和列分别进行一维DCT实现。norm='ortho' 使用正交归一化,保证变换的能量守恒。

DCT压缩函数
def dct_compress(image, threshold, threshold_type='hard'):
    coeffs = dct2(image)
    
    if threshold_type == 'hard':
        coeffs_thresh = hard_threshold(coeffs, threshold)
    else:
        coeffs_thresh = soft_threshold(coeffs, threshold)
    
    reconstructed = idct2(coeffs_thresh)
    return reconstructed, coeffs_thresh
小波压缩函数
def wavelet_compress(image, threshold, wavelet='db4', level=3, threshold_type='hard'):
    coeffs = pywt.wavedec2(image, wavelet, level=level)
    
    coeffs_thresh = []
    for i, c in enumerate(coeffs):
        if i == 0:
            coeffs_thresh.append(c)  # 保留近似系数
        else:
            if threshold_type == 'hard':
                coeffs_thresh.append(tuple(hard_threshold(arr, threshold) for arr in c))
            else:
                coeffs_thresh.append(tuple(soft_threshold(arr, threshold) for arr in c))
    
    reconstructed = pywt.waverec2(coeffs_thresh, wavelet)
    reconstructed = reconstructed[:image.shape[0], :image.shape[1]]
    return reconstructed, coeffs_thresh

pywt.wavedec2 进行二维小波分解,返回近似系数和各层细节系数。db4 是Daubechies 4阶小波,level=3 表示分解3层。近似系数(低频)通常不进行阈值处理以保留图像的整体结构。

质量评估指标
def mse(original, reconstructed):
    return np.mean((original - reconstructed) ** 2)

def psnr(original, reconstructed):
    mse_val = mse(original, reconstructed)
    if mse_val == 0:
        return float('inf')
    max_pixel = np.max(original)
    return 10 * np.log10(max_pixel ** 2 / mse_val)

MSE(均方误差)衡量重建图像与原始图像的平均像素差异。PSNR(峰值信噪比)以分贝为单位,值越高表示质量越好,通常30dB以上被认为是可接受的质量。

压缩结果保存
def save_dct_coeffs(coeffs, filename):
    nonzero_mask = coeffs != 0
    nonzero_values = coeffs[nonzero_mask]
    nonzero_indices = np.argwhere(nonzero_mask)
    
    np.savez_compressed(filename,
                        values=nonzero_values,
                        indices=nonzero_indices,
                        shape=coeffs.shape)
    return os.path.getsize(filename)

只保存非零系数及其位置,使用 npz 压缩格式进一步减小文件大小。压缩率等于原始图像大小除以压缩后文件大小。

DCT与小波比较

DCT是JPEG标准的基础,对平滑区域压缩效果好,但在高压缩率下会产生块效应。小波变换是JPEG2000的基础,提供多分辨率分析,在边缘保持方面表现更好,高压缩率下产生的伪影更自然(模糊而非块状)。

通过绘制PSNR与非零系数比例的曲线,可以比较两种方法在不同压缩率下的质量表现。通常在相同的非零系数数量下,小波变换能获得略高的PSNR,特别是对于包含大量边缘的图像。

使用numexpr进行数值计算

问题背景

处理一个信号处理问题,其中二维数组表示一批信号:

X \in \mathbb{R}^{N \times T}

每一行是长度为 T 的信号。要应用频率相关的衰减和二次惩罚,模型为:

Y_{i,t} = a_i e^{-bt} + c X_{i,t}^2

其中 a \in \mathbb{R}^N 是每个信号的系数,b, c \in \mathbb{R} 是标量参数,t = 0, \ldots, T-1 是时间索引。

原始代码
import numpy as np
import numexpr as ne

N, T = 1024, 4096

X = np.random.randn(N, T)
a = np.linspace(1.0, 2.0, N)
b = 0.01
c = 0.1
t = np.arange(T)

Y_numpy = a * np.exp(-b * t) + c * X**2
Y = ne.evaluate("a * exp(-b * t) + c * X**2")
各数组的形状
print(f"X.shape: {X.shape}")  # (1024, 4096) = (N, T)
print(f"a.shape: {a.shape}")  # (1024,) = (N,)
print(f"t.shape: {t.shape}")  # (4096,) = (T,)

执行结果:

X.shape: (1024, 4096)
a.shape: (1024,)
t.shape: (4096,)
Y 的预期形状: (1024, 4096)

XN \times T 的信号矩阵,a 是长度为 N 的系数向量,t 是长度为 T 的时间向量。期望的输出 Y 应该与 X 形状相同,即 (N, T)

代码无法工作的原因
try:
    Y_numpy = a * np.exp(-b * t) + c * X**2
except Exception as e:
    print(f"NumPy 错误: {e}")

执行结果:

解释:
a.shape = (1024,) (N,)
t.shape = (4096,) (T,)
np.exp(-b * t).shape = (4096,) (T,)
a * np.exp(-b * t) 需要广播 (N,) * (T,)
这在 NumPy 中会报错,因为形状不兼容

问题在于广播规则。a 的形状是 (N,)(1024,)np.exp(-b * t) 的形状是 (T,)(4096,)。NumPy的广播规则要求从右向左对齐维度,且每个维度要么相等,要么其中一个为1。(1024,)(4096,) 无法广播,因为两个维度都不为1且不相等。

正确的代码
# 将 a 变成列向量 (N, 1)
a_col = a.reshape(-1, 1)  # (N, 1)
print(f"a_col.shape: {a_col.shape}")

# NumPy 正确版本
Y_numpy = a_col * np.exp(-b * t) + c * X**2
print(f"NumPy Y.shape: {Y_numpy.shape}")

# numexpr 正确版本
Y_ne = ne.evaluate("a_col * exp(-b * t) + c * X**2")
print(f"numexpr Y.shape: {Y_ne.shape}")

# 验证结果一致
print(f"结果一致: {np.allclose(Y_numpy, Y_ne)}")

执行结果:

a_col.shape: (1024, 1)
NumPy Y.shape: (1024, 4096)
numexpr Y.shape: (1024, 4096)
结果一致: True

a 重塑为 (N, 1) 后,广播变为 (N, 1) * (T,)(N, 1) * (1, T)(N, T)。现在两个数组可以正确广播,a_col 的每一行(只有一个元素)与 t 的每个元素相乘,得到 N \times T 的矩阵。

内存和时间比较
import time

a_col = a.reshape(-1, 1)

# NumPy 计时
start = time.time()
for _ in range(10):
    Y_numpy = a_col * np.exp(-b * t) + c * X**2
numpy_time = (time.time() - start) / 10
print(f"NumPy 平均时间: {numpy_time * 1000:.2f} ms")

# numexpr 计时
start = time.time()
for _ in range(10):
    Y_ne = ne.evaluate("a_col * exp(-b * t) + c * X**2")
ne_time = (time.time() - start) / 10
print(f"numexpr 平均时间: {ne_time * 1000:.2f} ms")

print(f"加速比: {numpy_time / ne_time:.2f}x")

numexpr通常比纯NumPy快2-10倍,具体取决于表达式复杂度和数据大小。

内存分析方面,XY 各占用约32MB(1024 \times 4096 \times 8 字节)。NumPy在计算过程中会创建多个临时数组:-b * t 产生 (T,) 形状的数组,np.exp(-b * t) 产生 (T,) 形状的数组,a_col * np.exp(-b * t) 产生 (N, T) 形状的大型临时数组,X**2 产生 (N, T) 形状的大型临时数组,c * X**2 产生 (N, T) 形状的大型临时数组。总共约有3-4个大型临时数组。

numexpr的优势在于:逐元素计算,避免创建大型临时数组;更好的缓存利用率;自动多线程并行。

广播行为
print("NumPy 广播:")
print(f"  a_col.shape = {a_col.shape} (N, 1)")
print(f"  t.shape = {t.shape} (T,)")
print(f"  广播后: (N, 1) * (T,) → (N, T)")

result_shape = ne.evaluate("a_col * exp(-b * t)").shape
print(f"  numexpr 结果形状: {result_shape}")

NumPy和numexpr都执行广播,规则相同。不同之处在于numexpr在内部逐块处理数据,即使广播产生大型结果数组,也不会一次性分配所有临时内存。

临时数组数量分析

NumPy表达式 a_col * np.exp(-b * t) + c * X**2 的计算步骤:

步骤 1: temp1 = -b * t           → shape (T,)
步骤 2: temp2 = np.exp(temp1)    → shape (T,)
步骤 3: temp3 = a_col * temp2    → shape (N, T)  ← 大型临时数组
步骤 4: temp4 = X**2             → shape (N, T)  ← 大型临时数组
步骤 5: temp5 = c * temp4        → shape (N, T)  ← 大型临时数组
步骤 6: Y = temp3 + temp5        → shape (N, T)  ← 结果

大型临时数组至少3个,总临时内存约为 3 \times 1024 \times 4096 \times 8 \approx 96 MB。

numexpr使用分块计算策略,每次只处理一小块数据(通常是CPU缓存大小级别),临时数组非常小,通常只有几KB到几MB。

计算循环的执行位置

NumPy的计算循环在C语言层面执行,通过BLAS/LAPACK库或内置的ufuncs实现。每个操作(exp、乘法、加法)单独循环遍历整个数组,导致多次遍历数据,缓存效率较低。

numexpr的计算循环也在C语言层面执行,但使用虚拟机(VM)解释编译后的表达式。关键区别是单次遍历数据,一次性计算所有操作。numexpr还自动使用OpenMP进行多线程并行。

# 显示 numexpr 线程信息
print(f"numexpr 线程数: {ne.nthreads}")

# 对比单线程和多线程
ne.set_num_threads(1)
start = time.time()
for _ in range(10):
    Y_single = ne.evaluate("a_col * exp(-b * t) + c * X**2")
single_time = (time.time() - start) / 10

ne.set_num_threads(ne.detect_number_of_cores())
start = time.time()
for _ in range(10):
    Y_multi = ne.evaluate("a_col * exp(-b * t) + c * X**2")
multi_time = (time.time() - start) / 10

print(f"单线程时间: {single_time * 1000:.2f} ms")
print(f"多线程时间: {multi_time * 1000:.2f} ms")
print(f"多线程加速比: {single_time / multi_time:.2f}x")

多线程通常能提供2-8倍的额外加速,取决于CPU核心数和内存带宽。numexpr在处理大型数组的复杂数值表达式时特别有效,因为它同时解决了内存带宽瓶颈和计算并行化问题。


评论