Python读取多个文件后读取速度骤降问题求助
问题
使用Python处理LAMMPS轨迹文件时,初始读取速度可达45k it/s,但运行一段时间后突然降至600it/s,所有文件大小相近,无法定位原因。相关代码如下:
读取文件函数
import numpy as np import numba def read_data(f,read): nskip = 9 for i in range(nskip): ##skip these lines line = f.readline() if i==3: n = int(line.split()[0]) #how many particles types = np.zeros(n) pos = np.zeros((n,2)) ang = np.zeros((n)) for i in range(n): line = f.readline() if read: types[i] = int(line.split()[0]) #particle data pos[i,0] = float(line.split()[2]) pos[i,1] = float(line.split()[3]) ang[i] = float(line.split()[5]) return types,pos,ang
数据处理循环
eta_list = [0.03 ,0.0625 ,0.125 ,0.25 ,0.5 ,1.0 ,2.0 ,4.0 ,8.0 ,16.0, 32.0, 64.0] files = [1,2,3,4,5,6,7,8,9,10,11,12,13,14,15,16,17,18,19,20] chi = 1.25 Nskip = 5500 Ntotal = 6000 N_surv = [] for eta in eta_list: #loop over parameter eta num_surv = 0 num = 0 for fnum in files: #each parameter has 20 files fname = '../all_test_eta{}_chi{}_{}.lammpstrj'.format(eta,chi,fnum) f = open(fname,mode='r') for i in tqdm(range(Nskip)): #TQDM to time the read_data call _,_,_ = read_data(f,read=0) #skip these frames -- equilibriation for i in range(Ntotal-Nskip): a,b,c = read_data(f,read=1) #work on this data num_surv += np.sum(a==3) num+=1 f.close() num_surv = num_surv /(len(files)*num) N_surv.append(num_surv)
问题分析与优化方案
核心原因
- 磁盘IO瓶颈:
readline()逐行读取依赖操作系统文件缓存,初始高速是因为缓存命中,缓存耗尽后触发物理磁盘读取,速度骤降。 - 冗余计算与内存开销:
read=0时仍执行split()操作,且每次调用read_data都重新分配数组,内存回收开销随运行时间累积。 - 低效的帧跳过逻辑:逐行跳过帧时重复执行不必要的字符串解析,浪费CPU资源。
针对性优化
1. 批量读取文件,消除磁盘IO波动
一次性读取整个文件到内存,后续所有操作在内存中完成,彻底避免频繁磁盘IO:
def read_data(lines, idx, read): nskip = 9 # 跳过帧头的9行 for i in range(nskip): line = lines[idx] idx += 1 if i == 3: n = int(line.split()[0]) # 读取n行粒子数据 end_idx = idx + n particle_lines = lines[idx:end_idx] idx = end_idx if not read: return idx, None, None, None types = np.zeros(n, dtype=np.int32) pos = np.zeros((n,2), dtype=np.float64) ang = np.zeros(n, dtype=np.float64) for i, line in enumerate(particle_lines): parts = line.split() types[i] = int(parts[0]) pos[i,0] = float(parts[2]) pos[i,1] = float(parts[3]) ang[i] = float(parts[5]) return idx, types, pos, ang # 处理文件时先一次性读入所有行 with open(fname, 'r') as f: lines = f.read().splitlines() idx = 0 # 跳过Nskip帧 for _ in tqdm(range(Nskip)): idx, _, _, _ = read_data(lines, idx, read=0) # 处理有效帧 for _ in range(Ntotal - Nskip): idx, a, b, c = read_data(lines, idx, read=1) num_surv += np.sum(a == 3) num += 1
2. 复用数组,减少内存分配开销
提前获取文件中最大粒子数,创建可复用的数组,避免每次调用read_data时重复分配内存:
# 获取单个文件的最大粒子数 def get_max_n(fname): with open(fname, 'r') as f: lines = f.read().splitlines() idx = 0 max_n = 0 while idx < len(lines): # 跳3行到含n的行 idx += 3 n = int(lines[idx].split()[0]) max_n = max(max_n, n) # 跳过剩余帧头+粒子行 idx += 6 + n return max_n # 提前创建复用数组 max_n = get_max_n('../all_test_eta0.03_chi1.25_1.lammpstrj') reuse_types = np.zeros(max_n, dtype=np.int32) reuse_pos = np.zeros((max_n,2), dtype=np.float64) reuse_ang = np.zeros(max_n, dtype=np.float64) # 修改read_data函数复用数组 def read_data(lines, idx, read, types_arr, pos_arr, ang_arr): nskip =9 for i in range(nskip): line = lines[idx] idx +=1 if i ==3: n = int(line.split()[0]) end_idx = idx +n particle_lines = lines[idx:end_idx] idx = end_idx if not read: return idx, None, None, None # 截取数组前n个元素复用 types = types_arr[:n] pos = pos_arr[:n] ang = ang_arr[:n] for i, line in enumerate(particle_lines): parts = line.split() types[i] = int(parts[0]) pos[i,0] = float(parts[2]) pos[i,1] = float(parts[3]) ang[i] = float(parts[5]) return idx, types, pos, ang # 使用时传入复用数组 idx, a, b, c = read_data(lines, idx, read=1, types_arr=reuse_types, pos_arr=reuse_pos, ang_arr=reuse_ang)
3. 快速跳过帧,避免冗余计算
直接计算跳过帧的总行数,一次性移动索引,跳过不必要的字符串解析:
def skip_frames(lines, idx, num_frames): for _ in range(num_frames): # 跳3行到含n的行 idx +=3 n = int(lines[idx].split()[0]) # 跳过剩余6行帧头 + n行粒子数据 idx += 6 + n return idx # 替换原跳过循环 idx = skip_frames(lines, 0, Nskip)
4. Numba加速CPU密集型循环(可选)
如果粒子解析是CPU瓶颈,用Numba编译解析循环(需提前分割所有行):
@numba.jit(nopython=True) def process_particles(particle_parts, types, pos, ang): for i in range(len(particle_parts)): parts = particle_parts[i] types[i] = int(parts[0]) pos[i,0] = float(parts[2]) pos[i,1] = float(parts[3]) ang[i] = float(parts[5]) # 提前分割所有行的字符串 with open(fname, 'r') as f: all_parts = [line.split() for line in f if line.strip()]
优化效果
优化后,磁盘IO次数从数万次降至1次,内存分配开销大幅减少,读取速度会稳定维持在内存级别,不会出现后期骤降的情况。
内容的提问来源于stack exchange,提问作者Physics_iyer
相关产品推荐
相关产品推荐

