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

Python中顺序yield调用的延迟问题及优化咨询

问题描述

我正在编写代码读取存储在H5文件中的多组pandas.DataFrame并遍历其行,目标是通过PyTorch的IterableDataset处理数据集,但该问题并非PyTorch专属。

由于从磁盘读取每个H5文件耗时较长,我实现了以下逻辑:

  • 代码从磁盘读取第一个文件,并异步预取第二个文件;
  • 第一个DataFrame读取完成后,通过yield from遍历其行;
  • 遍历完成后,异步读取下一个文件,并开始遍历之前预取的文件。

相关代码如下:

def _load_next( file_list, index, device, labels, variables):
    if index >= len(file_list):
        return None
        
    thedata=pd.read_hdf(file_list[index], 'df')
    labels=torch.Tensor( thedata[labels].values).to(device)
    variables=torch.Tensor( thedata[variables].values).to(device)     

    return index, (labels,variables)   

class datasets( IterableDataset ):
    def __init__( self, path, device, variables, labels):
        self.files=glob.glob(path)
        self.device=device
        self.variables=variables
        self.labels=labels
        self.restart()
        
    def restart(self):
        print("Re-starting iterator")
        # read first file and submit prefetching of the following
        self.file_index, self.current_data=_load_next(self.files,0, self.device)
        self.prefetch=self.executor.submit(_load_next, self.files, self.file_index+1, self.device)   
        
    def __iter__(self):
       while True:
            yield from zip(self.current_data[0], self.current_data[1])
            result=self.prefetch.result()
            if result is None: 
                self.executor.shutdown(wait=False)
                raise StopIteration
            else:
                self.file_index, self.current_data = result
                self.prefetch=self.executor.submit(_load_next, self.files, self.file_index+1, self.device)

该逻辑运行正常,但每次yield from调用会耗时数秒,引入了不必要的延迟(延迟时长甚至超过预取下一个文件的时间)。请问是否可以消除该延迟,比如异步执行yield from?也欢迎其他优化思路,感谢!

优化方案

yield from本身是同步生成器操作,无法直接异步执行,但可以通过拆分数据批次、优化数据处理流程、扩容预取队列等方式消除延迟,以下是具体思路:

1. 拆分数据为小批次生成

当前yield from直接遍历整个DataFrame的行,单次生成大量数据项导致阻塞时间过长。可以将数据拆分为小批次,分散生成压力,同时让预取操作与批次生成并行:

def __iter__(self):
    batch_size = 64  # 根据设备内存/性能调整
    while True:
        labels, vars = self.current_data
        # 按批次拆分并生成数据
        for i in range(0, len(labels), batch_size):
            batch_labels = labels[i:i+batch_size]
            batch_vars = vars[i:i+batch_size]
            yield from zip(batch_labels, batch_vars)
        # 批次遍历完成后再切换预取文件
        result = self.prefetch.result()
        if result is None:
            self.executor.shutdown(wait=False)
            raise StopIteration
        else:
            self.file_index, self.current_data = result
            self.prefetch = self.executor.submit(_load_next, self.files, self.file_index+1, self.device)

这种方式缩短了单次yield的阻塞窗口,预取操作可以在批次生成的间隙完成,实现真正的并行。

2. 优化数据转换与设备迁移流程

当前_load_next中一次性将整个DataFrame转换为Tensor并迁移到设备,这一步可能占用大量时间。可以做两点优化:

  • 使用non_blocking=True开启非阻塞设备迁移,让数据传输与后续计算并行;
  • 将大尺寸数据转换拆分为小批次操作,分散到每个yield阶段。

修改后的_load_next和生成逻辑示例:

def _load_next(file_list, index, labels_col, variables_col):
    if index >= len(file_list):
        return None
    thedata = pd.read_hdf(file_list[index], 'df')
    # 只返回原始DataFrame片段,不做Tensor转换
    return index, (thedata[labels_col], thedata[variables_col])

def __iter__(self):
    batch_size = 64
    while True:
        labels_df, vars_df = self.current_data
        for i in range(0, len(labels_df), batch_size):
            # 批次转换并非阻塞迁移到设备
            batch_labels = torch.Tensor(labels_df.iloc[i:i+batch_size].values).to(self.device, non_blocking=True)
            batch_vars = torch.Tensor(vars_df.iloc[i:i+batch_size].values).to(self.device, non_blocking=True)
            yield from zip(batch_labels, batch_vars)
        # 后续切换文件逻辑不变
        result = self.prefetch.result()
        if result is None:
            self.executor.shutdown(wait=False)
            raise StopIteration
        else:
            self.file_index, self.current_data = result
            self.prefetch = self.executor.submit(_load_next, self.files, self.file_index+1, self.labels, self.variables)

3. 扩容预取队列

当前仅预取1个文件,若当前文件遍历耗时过长,可能出现等待间隙。可以用队列维护多个预取文件,确保始终有备用数据:

from queue import Queue
import concurrent.futures

class datasets(IterableDataset):
    def __init__(self, path, device, variables, labels, prefetch_count=2):
        self.files = glob.glob(path)
        self.device = device
        self.variables = variables
        self.labels = labels
        self.prefetch_count = prefetch_count
        self.queue = Queue(maxsize=prefetch_count)
        self.executor = concurrent.futures.ThreadPoolExecutor(max_workers=prefetch_count)
        self._fill_queue()
        
    def _fill_queue(self):
        # 初始化填充预取队列
        for idx in range(min(self.prefetch_count, len(self.files))):
            future = self.executor.submit(_load_next, self.files, idx, self.labels, self.variables)
            self.queue.put(future)
        self.next_index = self.prefetch_count
        
    def __iter__(self):
        while not self.queue.empty():
            future = self.queue.get()
            result = future.result()
            if result is None:
                break
            file_idx, (labels_df, vars_df) = result
            # 批次生成数据
            batch_size = 64
            for i in range(0, len(labels_df), batch_size):
                batch_labels = torch.Tensor(labels_df.iloc[i:i+batch_size].values).to(self.device, non_blocking=True)
                batch_vars = torch.Tensor(vars_df.iloc[i:i+batch_size].values).to(self.device, non_blocking=True)
                yield from zip(batch_labels, batch_vars)
            # 继续预取下一个文件
            if self.next_index < len(self.files):
                future = self.executor.submit(_load_next, self.files, self.next_index, self.labels, self.variables)
                self.queue.put(future)
                self.next_index += 1
        self.executor.shutdown(wait=False)

4. 底层读取优化

  • 使用pandas.read_hdf的chunksize参数分块读取H5文件,避免一次性加载整个DataFrame到内存,比如pd.read_hdf(file_path, 'df', chunksize=1024);
  • 若H5文件为自行生成,可优化存储格式:使用固定数据类型、选择合适的压缩级别,提升读取效率。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.05 14:15:18