云计算百科
云计算领域专业知识百科平台

Python-numpy-高级技巧

NumPy 高级技巧

🧠 向量化思维

NumPy 的精髓不是用 ndarray 存数据,而是用向量化操作代替循环。

反模式 vs 最佳实践

import numpy as np

data = np.random.rand(1000000)

# ❌ 循环(Python 级别,极慢)
result = np.empty_like(data)
for i in range(len(data)):
result[i] = np.sin(data[i]) if data[i] > 0.5 else np.cos(data[i])

# ✅ 向量化(C 级别)
mask = data > 0.5
result = np.where(mask, np.sin(data), np.cos(data))

# 速度差 10~100 倍

更多例子

# ❌ 逐元素处理
squared = np.array([x**2 for x in data])

# ✅
squared = data ** 2

# ❌ 手动计算行均值
means = np.zeros(len(mat))
for i in range(len(mat)):
means[i] = mat[i].mean()

# ✅
means = mat.mean(axis=1)


🎯 np.einsum — 爱因斯坦求和约定 ⭐

einsum 是最强大却常被忽视的函数。它用简洁的字符串表达式描述任意张量运算。

基本语法

# np.einsum(subscripts, *operands)
# subscripts: '输入1轴标签, 输入2轴标签 -> 输出轴标签'

经典运算

a = np.random.rand(3, 4)
b = np.random.rand(4, 5)
v = np.random.rand(3, 4)
w = np.random.rand(4)

# 矩阵乘法
C = np.einsum('ij,jk->ik', a, b) # = a @ b

# 逐元素乘法
C = np.einsum('ij,ij->ij', a, v) # = a * v

# 点积
s = np.einsum('i,i->', w, w) # = w @ w

# 外积
C = np.einsum('i,j->ij', w, w) # = np.outer(w, w)

# 转置
C = np.einsum('ij->ji', a) # = a.T

# 对角线
d = np.einsum('ii->i', a) # = np.diag(a)(取 3×4 的主对角线)

# 迹
t = np.einsum('ii->', a) # = np.trace(a)

# 按行求和
row_sum = np.einsum('ij->i', a) # = a.sum(axis=1)

# 按列求和
col_sum = np.einsum('ij->j', a) # = a.sum(axis=0)

# 所有元素求和
total = np.einsum('ij->', a) # = a.sum()

高级 einsum

# 批量矩阵乘法
A = np.random.rand(10, 3, 4) # batch=10
B = np.random.rand(10, 4, 5)
C = np.einsum('bij,bjk->bik', A, B) # (10, 3, 5)

# 注意力分数(QK^T / sqrt(d))
Q = np.random.rand(32, 100, 64) # (batch, seq, d_k)
K = np.random.rand(32, 100, 64)
scores = np.einsum('bqd,bkd->bqk', Q, K) / np.sqrt(64)

# 张量收缩(Tensor Contraction)
X = np.random.rand(2, 3, 4, 5)
Y = np.random.rand(5, 6)
Z = np.einsum('abcd,de->abce', X, Y) # (2, 3, 4, 6)

省略号 … 在 einsum 中

# 处理任意维度的张量
X = np.random.rand(2, 3, 4, 5)

# 对最后两维做矩阵乘法
Z = np.einsum('…ij,…jk->…ik', X, X) # (2, 3, 4, 4)


🏗️ 结构化数组(Structured Arrays)

存储异构数据(类似数据库表的行),每列可不同类型。

# 定义结构化 dtype
dtype = np.dtype([
('name', 'U20'), # Unicode 字符串,最长 20 字符
('age', 'i4'), # 32 位整数
('score', 'f8'), # 64 位浮点数
('passed', '?') # 布尔值
])

# 创建结构化数组
students = np.array([
('张三', 20, 92.5, True),
('李四', 22, 78.0, True),
('王五', 21, 55.0, False),
('赵六', 23, 88.5, True),
], dtype=dtype)

# 访问
print(students['name']) # ['张三' '李四' '王五' '赵六']
print(students['age'].mean()) # 21.5
print(students[students['passed']]) # 只取通过的学生
print(students[students['score'] > 60]['name']) # 及格的姓名

# 排序
sorted_by_score = np.sort(students, order='score')[::1]
print(sorted_by_score['name']) # 按分数降序

嵌套 dtype

# 嵌套结构
dtype = np.dtype([
('id', 'i4'),
('position', [
('x', 'f8'),
('y', 'f8'),
('z', 'f8')
]),
('color', [
('r', 'u1'),
('g', 'u1'),
('b', 'u1')
])
])

points = np.zeros(100, dtype=dtype)
points['position']['x'] = np.random.rand(100)
points['position']['y'] = np.random.rand(100)


🎭 掩码数组(Masked Arrays)

处理含缺失值 / 无效值的数据。

import numpy.ma as ma

# 创建掩码数组
data = np.array([1, 2, 999, 4, 999, 6])
masked = ma.masked_where(data == 999, data)
# 或: masked = ma.masked_values(data, -999)

print(masked)
# [1 2 — 4 — 6]

# 统计自动忽略掩码值
print(masked.mean()) # 3.25(而非包括 -999)
print(masked.std()) # 忽略掩码的标准差
print(masked.sum())

# 填充掩码值
filled = masked.filled(0) # 掩码替换为 0
print(filled) # [1 2 0 4 0 6]

# 创建掩码
mask = np.array([False, False, True, False, False, False])
masked = ma.array(data, mask=mask)

# 覆盖掩码
masked[0] = ma.masked # 手动标记为缺失

二维掩码数组

# 以条件创建掩码
data = np.random.randn(5, 5)
data[0, 0] = np.nan
data[2, 3] = np.nan

# 自动屏蔽 NaN
masked = ma.masked_invalid(data)

print(masked.mean()) # 忽略 NaN
print(masked.mean(axis=0)) # 每列均值(忽略 NaN)

# 获取掩码矩阵
print(masked.mask) # True = 被屏蔽


🔬 Strides — 理解内存布局

arr = np.arange(12).reshape(3, 4).astype(np.int64)
print(arr.strides) # (32, 8)
# axis=0: 跨 32 字节 = 4 个 int64
# axis=1: 跨 8 字节 = 1 个 int64

# C-order (row-major): 最后一维步长最小
# [[ 0 1 2 3]
# [ 4 5 6 7]
# [ 8 9 10 11]]
# 内存中: [0 1 2 3 4 5 6 7 8 9 10 11] ← 一行一行

# Fortran-order (column-major): 第一维步长最小
arr_f = np.asfortranarray(np.arange(12).reshape(3, 4).astype(np.int64))
print(arr_f.strides) # (8, 24)
# 内存中: [0 4 8 1 5 9 2 6 10 3 7 11] ← 一列一列

利用 strides 做技巧性操作

# 创建步进数组(无需复制)
from numpy.lib.stride_tricks import as_strided, sliding_window_view

# 滑动窗口(NumPy 1.20+)
arr = np.arange(10)
windows = np.lib.stride_tricks.sliding_window_view(arr, window_shape=3)
print(windows.shape) # (8, 3)
# [[0 1 2]
# [1 2 3]
# …
# [7 8 9]]

# 应用于二维
mat = np.arange(25).reshape(5, 5)
patches = np.lib.stride_tricks.sliding_window_view(mat, (2, 2))
print(patches.shape) # (4, 4, 2, 2) — 4×4 个 2×2 的块


⚡ 性能优化清单

1. 避免 Python 循环

# ❌
result = np.array([np.mean(mat[i, :]) for i in range(mat.shape[0])])

# ✅
result = mat.mean(axis=1)

2. 使用 ufunc 代替自定义函数

# ❌
result = np.array([1 / x if x != 0 else 0 for x in data])

# ✅
result = np.divide(1, data, where=data != 0, out=np.zeros_like(data))

3. 减少不必要的数据复制

# ❌ 链式操作每步创建临时数组
result = (arr * 2 + 1).mean()

# ✅ numexpr / numba 可减少中间体(或者直接接受,影响不大)
# NumPy 2.0+ 优化了这个,但仍需注意

4. 用 out 参数省内存

result = np.empty_like(a)
np.multiply(a, b, out=result) # 结果直接写入 result

5. 批量操作代替小操作

# ❌ 逐行处理
for i in range(mat.shape[0]):
mat[i] = mat[i] / mat[i].sum()

# ✅ 批量操作
mat = mat / mat.sum(axis=1, keepdims=True)

6. 数据对齐(内存连续性)

# 检查内存是否连续
print(arr.flags['C_CONTIGUOUS']) # C 连续
print(arr.flags['F_CONTIGUOUS']) # Fortran 连续

# 如果不连续,某些操作可能变慢
arr = np.ascontiguousarray(arr) # 强制 C 连续
arr = np.asfortranarray(arr) # 强制 Fortran 连续


🔧 自定义 ufunc

# 用 frompyfunc 创建 ufunc(比 Python 循环快,但仍不如内置 ufunc)
def my_func(x):
return x**2 + 3*x + 1

my_ufunc = np.frompyfunc(my_func, nin=1, nout=1)
result = my_ufunc(np.arange(10))
print(result.dtype) # object(frompyfunc 总是返回 object 类型)

# 强转类型
result = result.astype(np.float64)

# 更好的做法是用 vectorize
@np.vectorize
def my_vec_func(x):
return x**2 + 3*x + 1

result = my_vec_func(np.arange(10)).astype(np.float64)

⚠️ frompyfunc 和 vectorize 本质仍是 Python 循环,只是语法糖。追求极致性能应直接用 NumPy 内置 ufunc 或 Numba。


💡 实用小技巧

找到满足条件的第一个/最后一个索引

arr = np.array([0, 0, 1, 2, 3, 0, 4])

# 第一个 >1 的索引
first = np.argmax(arr > 1) # 3
# 最后一个 ≠0 的索引
last = len(arr) 1 np.argmax((arr != 0)[::1]) # 6
# 或用 where
indices = np.where(arr > 1)[0]
first, last = indices[0], indices[1]

分箱统计

# bincount: 整数分箱计数(比 np.histogram 快)
data = np.random.randint(0, 10, 100000)
counts = np.bincount(data) # [0-9 的出现次数]
print(counts) # [10001 10023 9945 …]

# bincount + weights: 按箱求和
values = np.random.rand(100000)
sums = np.bincount(data, weights=values)
means = sums / counts # 每类的均值

唯一值映射(替代 pandas factorize)

labels = np.array(['cat', 'dog', 'bird', 'cat', 'dog', 'fish'])

# 返回唯一值和编码
unique, inverse = np.unique(labels, return_inverse=True)
print(unique) # ['bird' 'cat' 'dog' 'fish'](已排序)
print(inverse) # [1 2 0 1 2 3] — 每个元素在 unique 中的索引

多维数组的 argpartition(Top-K)

arr = np.array([3, 1, 4, 1, 5, 9, 2, 6])

# 找前 3 大的值(不保证排序,但比 argsort 快得多)
indices = np.argpartition(arr, 3)[3:]
print(arr[indices]) # [5 9 6](不一定有序)

# 确保有序
top_k_indices = np.argpartition(arr, 3)[3:]
top_k_indices = top_k_indices[np.argsort(arr[top_k_indices])[::1]]
print(arr[top_k_indices]) # [9 6 5](已排序)

数组的属性赋值

arr = np.array([1, 2, 3, 4, 5])

# 对筛选出的元素批量赋值
arr[arr % 2 == 0] = 1 # [ 1 -1 3 -1 5]

# 用花式索引赋值
arr[[0, 2, 4]] = [10, 30, 50] # [10 -1 30 -1 50]


📊 与 Numba 结合(超高加速)

# pip install numba
from numba import jit
import numpy as np

# 对无法完全向量化的复杂算法,用 Numba JIT 编译
@jit(nopython=True)
def complex_algorithm(data):
result = np.empty_like(data)
for i in range(len(data)):
s = 0.0
for j in range(i+1):
s += data[j] * np.exp((ij) / 10) # 指数衰减加权
result[i] = s
return result

# 速度接近 C 语言
data = np.random.rand(100000)
result = complex_algorithm(data)


📝 速查表

需求代码
矩阵乘法 np.einsum('ij,jk->ik', A, B)
转置 np.einsum('ij->ji', A)
对角线 np.einsum('ii->i', A)
外积 np.einsum('i,j->ij', a, b)
批矩阵乘法 np.einsum('bij,bjk->bik', A, B)
结构化数组 np.array(data, dtype=[('name','U10'),('age','i4')])
掩码数组 ma.masked_where(cond, arr)
忽略 NaN ma.masked_invalid(arr)
滑动窗口 np.lib.stride_tricks.sliding_window_view(arr, k)
连续化 np.ascontiguousarray(arr)
找到非零 np.argmax(cond)
整数分箱 np.bincount(data)
唯一值编码 np.unique(labels, return_inverse=True)
Top-K np.argpartition(arr, -k)[-k:]

[[numpy-总览|← 返回总览]]

赞(0)
未经允许不得转载:网硕互联帮助中心 » Python-numpy-高级技巧
分享到: 更多 (0)

评论 抢沙发

评论前必须登录!