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

Python简单循环内存占用过高问题求助

天文数据互相关分析代码内存优化建议

我编写了如下Python代码,用于读取一组小型观测数据,执行互相关计算并保存图表。虽然已经添加了内存清理的代码,但运行数分钟后内存占用会在30-55GB间波动,Mac变得卡顿。即使仅读取数据子集(完整文件约6GB),内存问题仍未缓解。

import matplotlib.pyplot as plt
import numpy as np
import astropy.units as u
from sunkit_image.time_lag import cross_correlation, get_lags, max_cross_correlation, time_lag

time=np.linspace(0,43200,num=int(43200/12))
timeu = time * u.s

for i in range(len(folders)):             # loop over all dates
    os.chdir('/Volumes/LaCie/timelags/RARs/'+folders[i])
    print(folders[i])
    for j in range(len(pairs)):           # iterates over every pair of data sets
        for x in range(36):               # sets up a sliding 2-hour window that shifts 20 min at a time
            ch_a = np.load('dc'+pairs[j][0]+'.npy',allow_pickle=True)[()][100*x:(100*x)+600,:,:] # read in only necessary data (but entire file is only ~6 Gb)
            ch_b = np.load('dc'+pairs[j][1]+'.npy',allow_pickle=True)[()][100*x:(100*x)+600,:,:] # read in only necessary data (but entire file is only ~6 Gb)
            
            ctime= timeu[100*x:(100*x)+600] # sets up the correct time array
            print('ctime range:',ctime[0],ctime[-1],len(ctime))
            
            max_cc_map = max_cross_correlation(ch_a, ch_b, ctime)
            tl_map = time_lag(ch_a, ch_b, ctime)
            del ch_a # trying to deal with memory issue
            del ch_b # trying to deal with memory issue
            
            plt.close('all') # making sure I don't just create endless open plots
            fig = plt.figure()
            ax = fig.add_subplot()
            im = ax.imshow(np.flip(tl_map,axis=0), cmap="cubehelix", vmin=-6000, vmax=6000)
            cax = make_axes_locatable(ax).append_axes("right", size="5%", pad="10%")
            fig.colorbar(im, cax=cax,label=r"$\tau_{AB}$ [s]")
            plt.tight_layout()
            fig.savefig('timelag_'+pairs[j][0]+'_'+pairs[j][1]+'_'+str(x)+'.png',dpi=400)
            
            fig = plt.figure()
            ax = fig.add_subplot()
            im = ax.imshow(np.flip(max_cc_map,axis=0), cmap="plasma",vmin=0,vmax=1)
            cax = make_axes_locatable(ax).append_axes("right", size="5%", pad="10%")
            fig.colorbar(im, cax=cax,label=r"Max Cross-correlation")
            plt.tight_layout()
            fig.savefig('maxcc_'+pairs[j][0]+'_'+pairs[j][1]+'_'+str(x)+'.png',dpi=400)
            
            fig=plt.figure(figsize=(10,6))
            values_tl, bins_tl, bars = plt.hist(np.ravel(np.asarray(tl_map)),bins=np.arange(-6000,6000,12000/50),log=True,label='Time Lags')

            values_masked, bins_masked, bars = plt.hist(np.ravel(np.asarray(tl_map)[np.where(np.asarray(max_cc_map) > 0.25)])
                                          ,bins=np.arange(-6000,6000,12000/50),log=True,label='Masked CC > 0.25')

            values_masked2, bins_masked2, bars = plt.hist(np.ravel(np.asarray(tl_map)[np.where(np.asarray(max_cc_map) > 0.5)])
                                          ,bins=np.arange(-6000,6000,12000/50),log=True,label='Masked CC > 0.5')
            values_masked3, bins_masked3, bars = plt.hist(np.ravel(np.asarray(tl_map)[np.where(np.asarray(max_cc_map) > 0.75)])
                                          ,bins=np.arange(-6000,6000,12000/50),log=True,label='Masked CC > 0.75')

            plt.ylabel('Pixel Occurrence')
            plt.legend()
            fig.savefig('hist_tl_cc_'+pairs[j][0]+'_'+pairs[j][1]+'_'+str(x)+'.png',dpi=400)

优化建议

  • 补全缺失导入:代码使用了os.chdir()和make_axes_locatable但未导入对应模块,补上后避免隐式依赖问题:

    import os
    from mpl_toolkits.axes_grid1 import make_axes_locatable
    
  • 手动触发垃圾回收:del语句标记对象为可回收,但Python不会立刻释放内存,手动调用垃圾回收强制释放:

    import gc
    # 在del ch_a, ch_b后添加
    gc.collect()
    
  • 复用绘图对象:每次循环新建fig和ax会积累大量内存对象,改为提前创建对象,循环时清空轴内容复用:

    # 移到内层循环外
    fig_tl, ax_tl = plt.subplots()
    fig_cc, ax_cc = plt.subplots()
    fig_hist, ax_hist = plt.subplots(figsize=(10,6))
    
    # 内层循环中
    ax_tl.clear()
    # 绘制timelag图...
    
  • 减少数组冗余操作:

    • tl_map和max_cc_map本身已是numpy数组,无需重复调用np.asarray()
    • 用tl_map.flatten()替代np.ravel(np.asarray(tl_map)),避免临时数组创建
    • 提前计算直方图的bins,避免内层循环重复计算:
      # 放在pairs循环外
      hist_bins = np.arange(-6000, 6000, 12000/50)
      # 内层循环直接使用hist_bins
      
  • 替代数组翻转操作:绘图时通过origin='lower'参数替代np.flip(),无需创建翻转后的临时数组:

    im = ax.imshow(tl_map, cmap="cubehelix", vmin=-6000, vmax=6000, origin='lower')
    
  • 切换非交互式绘图后端:交互式后端会残留资源,代码开头添加:

    plt.switch_backend('Agg')
    
  • 检查核心计算函数的内存占用:max_cross_correlation和time_lag可能在内部创建大临时数组,若支持分块处理,可尝试按像素块拆分计算,避免一次性加载所有数据到内存。

内容的提问来源于stack exchange,提问作者weaselskinghenry

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.31 15:50:31