基于Numpy的自定义DataLoader处理高频Tick数据时慢且不稳定问题排查
问题描述
我正在开发类似DeepLOB的模型,使用高频Tick级金融数据。由于数据量庞大且需构造成时间序列格式,无法一次性将整个数据集加载到内存中,因此实现了一个从.npy文件批量读取数据的自定义DataLoader:
import numpy as np import os import gc class Np_DataLoader_Cache: def __init__(self, data_dir, file_list, batch_size, train_ratio=0.8): self.data_dir = data_dir self.batch_size = batch_size self.file_list = file_list self.train_ratio = train_ratio self.current_file_id = -1 self.current_data = None self.shapes = [] for file in file_list: with open(os.path.join(data_dir, file), 'rb') as f: version = np.lib.format.read_magic(f) shape, fortran_order, dtype = np.lib.format.read_array_header_1_0(f) self.shapes.append(shape[0]) self.prefix_sums = [0] for s in self.shapes: self.prefix_sums.append(self.prefix_sums[-1] + s) self.total_samples = self.prefix_sums[-1] self.train_end_idx = int(self.total_samples * self.train_ratio) self.valid_start_idx = self.train_end_idx def _load_file_if_needed(self, file_id): if self.current_file_id != file_id: if self.current_data is not None: del self.current_data gc.collect() filename = self.file_list[file_id] self.current_data = np.load(os.path.join(self.data_dir, filename), mmap_mode='r') self.current_file_id = file_id def _get_batch(self, start_row_global, end_row_global): files_needed = [] current_row = start_row_global file_id = 0 while current_row < end_row_global: while file_id < len(self.prefix_sums) - 1 and self.prefix_sums[file_id + 1] <= current_row: file_id += 1 self._load_file_if_needed(file_id) file_start_in_global = self.prefix_sums[file_id] file_start_in_file = current_row - file_start_in_global file_end_in_global = min(end_row_global, self.prefix_sums[file_id + 1]) file_end_in_file = file_end_in_global - file_start_in_global slice_data = self.current_data[file_start_in_file:file_end_in_file] files_needed.append(slice_data) current_row += slice_data.shape[0] return np.concatenate(files_needed, axis=0) def get_train_batch(self, batch_index): start_row = batch_index * self.batch_size end_row = min((batch_index + 1) * self.batch_size, self.train_end_idx) if start_row >= end_row: return None return self._get_batch(start_row, end_row) def get_valid_batch(self, batch_index): start_row = self.valid_start_idx + batch_index * self.batch_size end_row = min(self.valid_start_idx + (batch_index + 1) * self.batch_size, self.total_samples) if start_row >= end_row: return None return self._get_batch(start_row, end_row) def close(self): if self.current_data is not None: del self.current_data self.current_data = None self.current_file_id = -1 gc.collect()
训练过程中发现数据加载极慢,并非模型计算耗时,而是遍历数据集的时间。为此编写了批量加载速度测试代码:
import os data_dir = './LOB_OFI_sortcode_NoResample/' filelist = sorted([f for f in os.listdir(data_dir) if 'npy' in f])[:20] from DataLoader_NP_Cache import Np_DataLoader_Cache demo2 = Np_DataLoader_Cache(data_dir, filelist, batch_size = 4096, train_ratio = 0.8) batch_nums = demo2.train_end_idx // demo2.batch_size + 1 import time begin_t = time.time() very_first = time.time() for batch_index in range(batch_nums): mini_batch = demo2.get_train_batch(batch_index) if batch_index % 5000 == 0: end_t = time.time() print( f"Batch of {batch_index} Done. Process Ratio is {batch_index / batch_nums}" ) elapsed_time = end_t - begin_t print(f"This 5000Batchs using time: {elapsed_time:.2f} s") begin_t = time.time() end_t = time.time() elapsed_time = end_t - very_first print(f"Total: {elapsed_time:.2f} s")
测试结果显示:部分批次加载速度正常(约15秒/5000批次),但偶尔会有批次突然耗时100+秒,速度波动极大。卡顿出现时机随机,有时恢复,有时不恢复,常出现在第3或第4个5000批次阶段。怀疑是DataLoader实现或numpy的mmap_mode使用存在问题,寻求技术建议或解决方案。
解决方案与优化建议
优化mmap使用与文件缓存策略
- 移除主动调用
gc.collect()的逻辑:手动触发垃圾回收会带来不可预测的性能开销,尤其是在文件切换时。Python的垃圾回收机制会自动处理不再引用的内存映射对象,手动调用反而会打断正常的IO流程。 - 实现LRU缓存保留多个文件映射:当前每次仅缓存一个文件,跨文件批次或频繁切换文件时会重复加载。可以缓存2-3个最近访问的文件,减少重复加载次数,示例代码:
from collections import OrderedDict def __init__(self, data_dir, file_list, batch_size, train_ratio=0.8, cache_size=2): # 保留原有初始化代码 self.cache = OrderedDict() self.cache_size = cache_size def _load_file_if_needed(self, file_id): if file_id not in self.cache: if len(self.cache) >= self.cache_size: self.cache.popitem(last=False) # 移除最早缓存的文件 filename = self.file_list[file_id] self.cache[file_id] = np.load(os.path.join(self.data_dir, filename), mmap_mode='r') self.current_file_id = file_id self.current_data = self.cache[file_id]
优化全局索引到文件的映射逻辑
- 使用二分查找替代线性遍历:当前
_get_batch中每次线性查找文件ID,当文件数量多时会累积开销。改用bisect模块实现二分查找,将复杂度从O(n)降到O(log n):import bisect # 在_get_batch中替换file_id查找逻辑 file_id = bisect.bisect_right(self.prefix_sums, current_row) - 1
减少跨文件批次的拼接开销
- 对齐批次与文件边界:如果可以重新组织.npy文件,让每个文件的样本数为
batch_size的整数倍,避免大部分批次跨文件,减少np.concatenate的使用。 - 避免不必要的数据复制:mmap切片是零拷贝操作,但
np.concatenate会生成新数组。如果模型支持,直接返回切片列表给模型处理,跳过拼接步骤。
IO层面的优化
- 更换高速存储介质:高频金融数据IO需求高,机械硬盘(HDD)换成SSD或NVMe磁盘能大幅降低随机IO延迟。
- 优化系统文件缓存:确保系统页缓存足够大,让常用.npy文件缓存在内存中,减少磁盘读取。可调整Linux系统的
vm.dirty_ratio等参数优化缓存策略,注意监控内存占用。 - 并行预加载数据:启动后台线程预加载下一个要访问的文件,让数据加载与模型计算重叠。使用
threading模块实现空闲时提前加载后续文件。
排查系统层面干扰
- 监控磁盘IO使用率:使用
iostat(Linux)或资源监视器(Windows)检查是否有后台进程(如备份、杀毒软件)占用IO资源,导致读取阻塞。 - 检查内存与swap使用:系统内存不足会触发swap,导致严重IO卡顿。确保可用内存足够容纳至少几个文件的映射数据,避免频繁触发swap。
内容的提问来源于stack exchange,提问作者ZaixinDong
相关产品推荐
相关产品推荐

