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

使用Matplotlib绘制十亿级数组折线图时内存耗尽求助

问题

在Google Colab中使用numpy内存映射数组创建了包含10亿条数据的数组,计算阶段一切正常,但调用plt.show()或plt.savefig()绘制折线图时内存突然耗尽。尝试过散点图的大内存数据绘图优化方案,但无法适配折线图场景,寻求解决办法。

原代码如下:

import numpy as np
import matplotlib.pyplot as plt
from tqdm import tqdm

width = 40
height = 8
transition = 0.9
chunk_size = int(1E6)  # Define a chunk size for processing data

n = int(1E8)

def calculate_array(n, tr, filename):
    a = np.memmap(filename, dtype='float64', mode='w+', shape=(n + 1,))
    if tr<1:
        for k in range(n, int(n * tr), -1):
            a[k] = 1 / (1 - tr) * (n - k) / n
    elif tr==1:
        a[n]=1

    for k in tqdm(range(int(n * tr), 0, -1), desc=f"Calculating array for t={tr}"):
        a[k] = 17 # of course in reality it's more complicated but let us cast aside those details
    
    a.flush()  # Ensure changes are written to disk

# Values of t to be used
t_values = [1, 0.99, 0.9, 0.8, 0.7, 0.6, 0.5]

# Create a memory-mapped file for x_values
x_filename = "x_values.dat"
x_values = np.memmap(x_filename, dtype='int64', mode='w+', shape=(n + 1,))
x_values[:] = np.arange(n + 1)
x_values.flush()

# Plot with both log axes
for t in t_values:
    plt.figure(figsize=(width, 8))
    filename = f"array_t_{t}.dat"
    calculate_array(n, t, filename)
    
    # Open the memory-mapped arrays for reading
    a = np.memmap(filename, dtype='float64', mode='r', shape=(n + 1,))
    x_values = np.memmap(x_filename, dtype='int64', mode='r', shape=(n + 1,))
    
    # Plot in chunks
    for start in tqdm(range(1, n, chunk_size), desc=f"Plotting array for t={t}"):
        end = min(start + chunk_size, n)
        x_chunk = x_values[start:end]
        y_chunk = a[start:end]
        
        plt.scatter(x_chunk, y_chunk, s=1, c='blue')
        
        # For the first chunk, connect the last point of the previous chunk
        if start > 1:
            plt.plot(x_values[start-1:end], a[start-1:end], linestyle='-', alpha=0.6, color='blue')
        else:
            plt.plot(x_chunk, y_chunk, linestyle='-', alpha=0.6, color='blue')
    
    plt.xscale('log')
    plt.xlabel('Index (log scale)')
    plt.yscale('symlog')
    plt.ylabel('a(k) (symlog scale)')
    plt.title(f'Individual Plot of the array a with both axes in log scale for t={t}')
    plt.legend([f't={t}'])
    plt.grid(True)
    plt.show()

解决方案

原代码内存耗尽的核心原因是:Matplotlib会将所有绘制的图形元素(每个chunk的折线、散点)都保存在内存中,10亿条数据会生成数百万个绘图对象,直接撑爆内存。针对折线图+对数轴的场景,优化思路如下:

优化要点

  1. 降采样而非逐chunk绘制:对数轴下,大部分密集点在视觉上完全重叠,无需绘制全部数据。可以按固定步长采样,或按对数区间均匀采样,大幅减少绘图元素数量。
  2. 直接利用内存映射数组切片采样:无需加载全部数据,通过numpy切片直接从内存映射文件中提取采样点,避免内存占用。
  3. 关闭交互式绘图直接保存:Colab中plt.show()会启动交互式绘图后端,额外占用内存,改为直接用plt.savefig()保存图片后关闭图形,释放资源。
  4. 及时清理图形对象:绘制完每个t对应的图后,用plt.close()清理当前图形,释放内存。

修改后的代码

import numpy as np
import matplotlib.pyplot as plt
from tqdm import tqdm

width = 40
height = 8
transition = 0.9
sample_rate = 1000  # 每1000个点取1个,可根据需求调整
# 可选对数均匀采样:生成对数区间的均匀索引
# sample_indices = np.unique(np.logspace(0, np.log10(n), num=10000, dtype=int))

n = int(1E8)

def calculate_array(n, tr, filename):
    a = np.memmap(filename, dtype='float64', mode='w+', shape=(n + 1,))
    if tr < 1:
        for k in range(n, int(n * tr), -1):
            a[k] = 1 / (1 - tr) * (n - k) / n
    elif tr == 1:
        a[n] = 1

    for k in tqdm(range(int(n * tr), 0, -1), desc=f"Calculating array for t={tr}"):
        a[k] = 17  # 保留实际计算逻辑
    
    a.flush()

# Values of t to be used
t_values = [1, 0.99, 0.9, 0.8, 0.7, 0.6, 0.5]

# Create a memory-mapped file for x_values
x_filename = "x_values.dat"
x_values = np.memmap(x_filename, dtype='int64', mode='w+', shape=(n + 1,))
x_values[:] = np.arange(n + 1)
x_values.flush()

# 切换到非交互式绘图后端,减少内存开销
plt.switch_backend('Agg')

for t in t_values:
    filename = f"array_t_{t}.dat"
    calculate_array(n, t, filename)
    
    # 打开内存映射数组
    a = np.memmap(filename, dtype='float64', mode='r', shape=(n + 1,))
    x_values = np.memmap(x_filename, dtype='int64', mode='r', shape=(n + 1,))
    
    # 降采样:按固定步长取点
    sample_indices = np.arange(1, n, sample_rate)
    x_sample = x_values[sample_indices]
    y_sample = a[sample_indices]
    
    # 绘制折线图(可选叠加采样后的散点)
    plt.figure(figsize=(width, height))
    plt.plot(x_sample, y_sample, linestyle='-', alpha=0.6, color='blue')
    # 若需要散点,仅绘制采样后的点
    # plt.scatter(x_sample, y_sample, s=1, c='blue')
    
    plt.xscale('log')
    plt.xlabel('Index (log scale)')
    plt.yscale('symlog')
    plt.ylabel('a(k) (symlog scale)')
    plt.title(f'Plot of array a (log axes) for t={t}')
    plt.legend([f't={t}'])
    plt.grid(True)
    
    # 直接保存图片,不调用show()
    plt.savefig(f"plot_t_{t}.png", bbox_inches='tight')
    plt.close()  # 关闭当前图形,释放内存
    print(f"Plot saved as plot_t_{t}.png")

额外说明

  • 若需要更精准的对数轴采样,可替换采样逻辑为对数均匀分布的索引(代码中注释部分),确保对数轴上每个区间的点数均匀,视觉效果更好。
  • 采样步长sample_rate可根据图形清晰度和内存情况调整:步长越大内存占用越低,步长越小细节越丰富。
  • Colab中使用Agg后端可避免交互式绘图的内存开销,如需预览可先用小批量数据测试,再用大数据运行保存。

内容的提问来源于stack exchange,提问作者D.R

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.21 17:02:33