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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.16 09:01:00