如何用datasets.Dataset.from_csv()合并指定列,或通过PyArrow实现?
问题解答
一、datasets读取阶段直接合并列?
不行。datasets.Dataset.from_csv 读取CSV时会自动将每一列解析为独立特征,没有内置参数支持读取时直接合并指定列。但读取后可以用高效的批量转换方式完成需求,不需要更换工具。
二、PyArrow读取时实现合并?
可以。PyArrow在读取CSV后(或读取过程中)支持灵活的列合并操作,属于读取阶段的处理:
import pyarrow as pa import pyarrow.csv as csv # 读取CSV为PyArrow Table table = csv.read_csv("your_file.csv") # 获取所有列名 cols = table.column_names # 将0-4列合并为feature1(数组类型列),剩余列合并为feature2 feature1 = table.select(cols[:5]).combine_chunks() feature2 = table.select(cols[5:]).combine_chunks() # 构造包含合并后列的新Table merged_table = pa.table({ "feature1": feature1, "feature2": feature2 }) # 如需转换为datasets.Dataset from datasets import Dataset dataset = Dataset(merged_table)
combine_chunks() 会将多列数据合并为一个数组类型的列,整个过程基于PyArrow的高效列式存储,性能优于逐行处理。
三、datasets读取后的高效转换方法
如果已经用datasets.Dataset.from_csv完成读取,推荐使用批量map操作实现合并,避免逐行处理带来的性能损耗:
from datasets import Dataset # 读取原始数据集 raw_dataset = Dataset.from_csv("your_file.csv") # 获取所有列名 all_columns = raw_dataset.column_names # 定义批量合并函数 def merge_columns(batch): # 合并0-4列:将多列的行数据转成每行的列表,打包为feature1 feature1 = [list(row) for row in zip(*[batch[col] for col in all_columns[:5]])] # 合并剩余列:同理生成feature2 feature2 = [list(row) for row in zip(*[batch[col] for col in all_columns[5:]])] # 返回新的batch结构 return {"feature1": feature1, "feature2": feature2} # 应用批量转换,移除原始列 merged_dataset = raw_dataset.map( merge_columns, batched=True, # 开启批量处理提升效率 remove_columns=all_columns )
这种方式利用datasets的批量处理能力,在底层基于PyArrow加速,处理大文件时效率远高于逐行遍历。
内容的提问来源于stack exchange,提问作者Wang
相关产品推荐
相关产品推荐

