TensorFlow新手:如何读取指定CSV文件并加载指定列数据
用TensorFlow读取多CSV文件并筛选指定列
你已经通过tf.data.Dataset.list_files获取了2018年所有目标文件的路径,接下来只需完成读取文件内容和筛选指定列两个核心步骤,具体实现如下:
1. 定义CSV解析函数
这个函数负责处理单个CSV文件:跳过表头、解析每行数据,并提取你指定索引的列:
import tensorflow as tf def parse_csv(file_path): # 按行读取文件内容 raw_lines = tf.data.TextLineDataset(file_path) # 跳过第一行表头 raw_lines = raw_lines.skip(1) # 定义各列的默认解析类型(需匹配数据集实际类型,这里假设共11列) defaults = [tf.string] * 11 # 解析每一行数据 parsed_rows = raw_lines.map(lambda line: tf.io.decode_csv(line, record_defaults=defaults)) # 提取指定索引的列:[1,3,6,8,10] def pick_target_cols(*cols): return (cols[1], cols[3], cols[6], cols[8], cols[10]) return parsed_rows.map(pick_target_cols)
2. 批量读取多文件并组合数据集
用interleave实现多文件并行读取,最后可根据需求设置批次大小:
# 列出所有2018年的CSV文件路径 data_2018_paths = tf.data.Dataset.list_files("./raw/*2018*") # 并行解析所有文件,组合成最终数据集 dataset = data_2018_paths.interleave( parse_csv, num_parallel_calls=tf.data.AUTOTUNE # 自动适配CPU并行能力 ) # 可选:设置批次大小(根据你的训练/分析需求调整) dataset = dataset.batch(32) # 测试读取一批数据验证结果 for batch in dataset.take(1): print("提取的列数据示例:") for idx, col in enumerate(batch): print(f"第{idx+1}列前5个样本:{col[:5]}")
补充优化提示
- 如果数据集列类型多样(如整数、浮点数),请修改
defaults中的对应类型(比如数值列用tf.int32/tf.float32),避免解析错误。 - 处理大文件时,建议添加
prefetch让数据加载与后续计算并行,提升效率:dataset = dataset.batch(32).prefetch(tf.data.AUTOTUNE)
内容的提问来源于stack exchange,提问作者Shawn Brar
相关产品推荐
相关产品推荐

