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

基于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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.12 17:30:53