如何用Numba加速二次映射迭代绘图代码?报错及优化咨询
解决Numba加速问题及代码优化方案
一、为什么加@njit会报错?
直接给func1加@njit必然报错,核心原因有三个:
- Numba不支持Python列表的动态
append操作:列表是动态扩容的动态类型容器,Numba需要静态类型结构才能完成编译优化,无法处理这种动态行为。 - Numba无法兼容matplotlib绘图代码:Numba仅专注于数值计算逻辑,对Python高层GUI/绘图类API完全不支持。
- 原代码
range((z/2),0,-1)存在语法问题:若z为整数,z/2会返回浮点数,而range仅接受整数参数,应改为整数除法z//2。
正确的Numba用法:拆分计算与绘图逻辑
把数值迭代的核心逻辑抽成独立函数用Numba加速,绘图等高层操作留在普通Python代码中:
from matplotlib import pyplot as plt import numpy as np from numba import njit # Numba加速核心计算逻辑 @njit def compute_points(x_start, y_start, a, iterations, skip=100): # 预分配numpy数组,避免动态append的内存开销 x_arr = np.empty(iterations - skip, dtype=np.float64) y_arr = np.empty(iterations - skip, dtype=np.float64) x = x_start y = y_start # 先跳过前100次迭代,不用存储无效点 for _ in range(skip): xnew = a[0] + a[1]*x + a[2]*x*x + a[3]*x*y + a[4]*y*y + a[5]*y ynew = a[6] + a[7]*x + a[8]*x*x + a[9]*x*y + a[10]*y*y + a[11]*y x, y = xnew, ynew # 存储有效迭代点 for i in range(iterations - skip): xnew = a[0] + a[1]*x + a[2]*x*x + a[3]*x*y + a[4]*y*y + a[5]*y ynew = a[6] + a[7]*x + a[8]*x*x + a[9]*x*y + a[10]*y*y + a[11]*y x, y = xnew, ynew x_arr[i] = x y_arr[i] = y return x_arr, y_arr def func1(z): x_start = 0.215 y_start = 0.512 # 用numpy数组存储参数a,比Python列表更适合数值计算 a = np.array([0.123,0.234,0.345,0.456,0.567,0.678,0.789,0.890,0.012,0.123,0.234,0.345], dtype=np.float64) # 绘图样式仅设置一次,避免重复操作 plt.style.use("dark_background") # 改用整数除法,避免浮点数传入range报错 for q in range(z//2, 0, -1): x_data, y_data = compute_points(x_start, y_start, a, 2500000) name = f"C:\\Users\\user\\imgs\\{q}.png" plt.scatter(x_data, y_data, s=0.001, marker='.', linewidth=0, c='#ffd769') plt.axis("off") plt.savefig(name, dpi=800) plt.clf() # numpy数组直接做增量操作,比列表推导式效率更高 a += 0.0002 if __name__=='__main__': func1(int(input("Enter Number of Frames (even): ")))
二、进一步优化加速的手段
1. 大幅提升绘图速度
matplotlib的scatter绘制250万点速度极慢,可改用密度图替代,速度提升几个数量级:
# 用histogram2d生成密度图代替scatter bins = 2000 # 可根据需求调整分辨率 counts, xedges, yedges = np.histogram2d(x_data, y_data, bins=bins) plt.imshow(counts.T, origin='lower', cmap='YlOrBr', extent=[xedges[0], xedges[-1], yedges[0], yedges[-1]])
这种方式不绘制单个点,而是渲染点的密度分布,视觉效果与散点图接近,但速度快很多。
2. 并行处理多帧
每个帧的计算与绘图完全独立,可使用多进程并行处理:
from multiprocessing import Pool def process_frame(q, z_half, a_initial, x_start, y_start): # 计算当前帧对应的a参数值 a = a_initial + 0.0002 * (z_half - q) x_data, y_data = compute_points(x_start, y_start, a, 2500000) plt.style.use("dark_background") plt.scatter(x_data, y_data, s=0.001, marker='.', linewidth=0, c='#ffd769') plt.axis("off") plt.savefig(f"C:\\Users\\user\\imgs\\{q}.png", dpi=800) plt.clf() if __name__=='__main__': z = int(input("Enter Number of Frames (even): ")) z_half = z//2 x_start = 0.215 y_start = 0.512 a_initial = np.array([0.123,0.234,0.345,0.456,0.567,0.678,0.789,0.890,0.012,0.123,0.234,0.345], dtype=np.float64) with Pool() as p: # 生成所有帧的处理参数 args = [(q, z_half, a_initial, x_start, y_start) for q in range(z_half, 0, -1)] p.starmap(process_frame, args)
注意:多进程模式下,需确保每个子进程独立初始化matplotlib环境。
3. 减少内存占用
如果不需要保留所有迭代历史,可在计算时直接跳过前100次迭代,仅存储有效点(已在前面的compute_points函数中实现),能减少约4%的内存占用,对大迭代次数场景更友好。
内容的提问来源于stack exchange,提问作者jdr
相关产品推荐
相关产品推荐

