如何在数据集无法全量加载时将PyArrow与PyTorch Dataset集成?
处理多Parquet文件的PyTorch Dataset实现(不加载全量数据)
你之前的Dataset是基于内存数据实现的,但如果数据是目录下的多个Parquet文件,要避免全量加载到内存,可以借助PyArrow实现按需加载的Dataset,核心思路是先记录每个Parquet文件的行数,通过全局索引快速定位到目标文件和行,再在__getitem__时读取对应数据。
实现步骤
预扫描Parquet文件,构建索引映射
遍历目标目录下的所有Parquet文件,用PyArrow读取每个文件的元数据获取行数(无需加载实际数据),快速统计全局总行数,同时记录每个文件对应的行范围(比如第0-1000行在file1.parquet,1001-2000行在file2.parquet)。实现按需加载的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
相关产品推荐
相关产品推荐

