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

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)

问题分析与优化方案

核心原因

  1. 磁盘IO瓶颈:readline()逐行读取依赖操作系统文件缓存,初始高速是因为缓存命中,缓存耗尽后触发物理磁盘读取,速度骤降。
  2. 冗余计算与内存开销:read=0时仍执行split()操作,且每次调用read_data都重新分配数组,内存回收开销随运行时间累积。
  3. 低效的帧跳过逻辑:逐行跳过帧时重复执行不必要的字符串解析,浪费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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.17 19:39:55