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

如何在数据集无法全量加载时将PyArrow与PyTorch Dataset集成?

处理多Parquet文件的PyTorch Dataset实现(不加载全量数据)

你之前的Dataset是基于内存数据实现的,但如果数据是目录下的多个Parquet文件,要避免全量加载到内存,可以借助PyArrow实现按需加载的Dataset,核心思路是先记录每个Parquet文件的行数,通过全局索引快速定位到目标文件和行,再在__getitem__时读取对应数据。

实现步骤

  1. 预扫描Parquet文件,构建索引映射
    遍历目标目录下的所有Parquet文件,用PyArrow读取每个文件的元数据获取行数(无需加载实际数据),快速统计全局总行数,同时记录每个文件对应的行范围(比如第0-1000行在file1.parquet,1001-2000行在file2.parquet)。

  2. 实现按需加载的Dataset类
    在__init__里完成文件扫描和索引映射;__len__返回全局总行数;__getitem__根据传入的全局索引,找到对应的Parquet文件和文件内的行号,再用PyArrow读取该行数据并转换为PyTorch张量返回。

完整代码示例

import torch
from torch.utils.data import Dataset
import pyarrow.parquet as pq
import os

class CustomParquetDataset(Dataset):
    def __init__(self, parquet_dir):
        self.parquet_dir = parquet_dir
        # 收集目录下所有Parquet文件路径
        self.file_paths = [os.path.join(parquet_dir, f) for f in os.listdir(parquet_dir) if f.endswith('.parquet')]
        # 存储每个文件的起始行索引、结束行索引和文件路径
        self.file_row_info = []
        total_rows = 0
        
        # 预扫描所有文件统计行数,构建行范围映射
        for file_path in self.file_paths:
            metadata = pq.read_metadata(file_path)
            num_rows = metadata.num_rows
            self.file_row_info.append((total_rows, total_rows + num_rows, file_path))
            total_rows += num_rows
        
        self.total_rows = total_rows

    def __len__(self):
        return self.total_rows

    def __getitem__(self, index):
        # 定位目标文件和文件内的行号
        target_file = None
        row_in_file = None
        for start, end, file_path in self.file_row_info:
            if start <= index < end:
                row_in_file = index - start
                target_file = file_path
                break
        
        # 读取指定行数据
        table = pq.read_table(target_file, rows=[row_in_file])
        df = table.to_pandas()
        # 提取特征和标签,转换为PyTorch张量(根据实际数据类型调整dtype)
        x = torch.tensor(df.iloc[0, :-1].values, dtype=torch.float32)
        y = torch.tensor(df.iloc[0, -1].values, dtype=torch.float32)
        
        return x, y

注意事项

  • 性能优化:如果单Parquet文件体积较大,read_table指定行的方式依然高效,因为PyArrow支持列存储的随机访问;也可以提前缓存每个文件的行索引范围,避免每次遍历查找。
  • 多进程适配:使用DataLoader时若设置num_workers>0,建议在__getitem__内每次读取文件时重新打开,避免跨进程共享文件对象导致的异常。
  • 类型适配:根据实际任务调整张量的dtype,比如分类任务的标签可改用torch.long类型。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.31 11:11:20