You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何用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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.18 08:23:12