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

TensorFlow处理大CSV文件:逐行读取与批量训练方案咨询

大体积CSV文件的Keras训练解决方案

1. 能否直接向model.fit传入文件?

不行。model.fit() 仅接受张量、NumPy数组、或实现迭代接口的数据集对象(如tf.data.Dataset、Keras Sequence),无法直接传入原始文件路径,否则会触发类型错误——框架无法自动解析未处理的文件内容。

2. 逐行读取文件并训练的两种方案

方案一:使用tf.data.Dataset(推荐,原生适配TensorFlow)

利用TensorFlow的tf.data.TextLineDataset逐行读取CSV,配合自定义解析函数处理每行数据,最后生成batch喂给模型。这种方式支持并行加载、预取,效率更高。

单文件示例(x和y同文件):

import tensorflow as tf

# 定义CSV解析函数,根据你的数据格式调整特征/标签的数量、类型
def parse_csv_line(line):
    num_features = 10  # 替换为你的特征列数量
    # 定义每个字段的默认类型(假设所有字段为浮点型)
    record_defaults = [tf.float32] * (num_features + 1)
    parsed = tf.io.decode_csv(line, record_defaults=record_defaults)
    features = tf.stack(parsed[:-1])
    label = parsed[-1]
    return features, label

# 加载CSV文件,跳过表头
dataset = tf.data.TextLineDataset("x_y_combined.csv").skip(1)
# 并行解析每行数据
dataset = dataset.map(parse_csv_line, num_parallel_calls=tf.data.AUTOTUNE)
# 设置batch size并预取数据提升训练效率
dataset = dataset.batch(32).prefetch(tf.data.AUTOTUNE)

# 喂给模型训练
model.fit(dataset, epochs=10)

分离文件示例(x和y分两个CSV):

def parse_x_line(line):
    record_defaults = [tf.float32] * 10  # 仅解析特征列
    return tf.stack(tf.io.decode_csv(line, record_defaults=record_defaults))

def parse_y_line(line):
    return tf.io.decode_csv(line, record_defaults=[tf.float32])[0]

x_dataset = tf.data.TextLineDataset("x.csv").skip(1).map(parse_x_line)
y_dataset = tf.data.TextLineDataset("y.csv").skip(1).map(parse_y_line)
# 合并特征与标签数据集
dataset = tf.data.Dataset.zip((x_dataset, y_dataset)).batch(32).prefetch(tf.data.AUTOTUNE)

model.fit(dataset, epochs=10)

方案二:自定义Keras Sequence类

如果更习惯用Pandas处理数据,可以继承keras.utils.Sequence实现自定义生成器,逐块读取CSV并返回batch。这种方式完全由Python控制,灵活性高。

import pandas as pd
import numpy as np
from tensorflow.keras.utils import Sequence

class CSVBatchGenerator(Sequence):
    def __init__(self, x_file, y_file, batch_size=32):
        self.x_file = x_file
        self.y_file = y_file
        self.batch_size = batch_size
        # 计算总样本数(减去表头行)
        with open(x_file) as f:
            self.total_samples = sum(1 for _ in f) - 1
        self.total_batches = (self.total_samples + batch_size - 1) // batch_size

    def __len__(self):
        return self.total_batches

    def __getitem__(self, idx):
        # 计算当前batch的起始/结束行(+1跳过表头)
        start_row = idx * self.batch_size + 1
        end_row = start_row + self.batch_size
        # 读取当前batch的特征与标签
        x_batch = pd.read_csv(self.x_file, skiprows=start_row, nrows=self.batch_size, header=None).values
        y_batch = pd.read_csv(self.y_file, skiprows=start_row, nrows=self.batch_size, header=None).values
        return x_batch, y_batch

# 初始化生成器并训练
generator = CSVBatchGenerator("x.csv", "y.csv", batch_size=32)
model.fit(generator, epochs=10)

3. 扩展到目录下多文件批量加载训练

方案一:tf.data.Dataset多文件加载(推荐)

用tf.data.Dataset.list_files获取目录下所有CSV文件,再用interleave并行读取每个文件,实现多文件批量加载:

import tensorflow as tf

def parse_csv_line(line):
    num_features = 10
    record_defaults = [tf.float32] * (num_features + 1)
    parsed = tf.io.decode_csv(line, record_defaults=record_defaults)
    return tf.stack(parsed[:-1]), parsed[-1]

# 获取目录下所有CSV文件路径
file_paths = tf.data.Dataset.list_files("./data_dir/*.csv")
# 并行读取每个文件,跳过表头后解析数据
dataset = file_paths.interleave(
    lambda path: tf.data.TextLineDataset(path).skip(1),
    num_parallel_calls=tf.data.AUTOTUNE
)
dataset = dataset.map(parse_csv_line, num_parallel_calls=tf.data.AUTOTUNE)
# 打乱数据、设置batch size并预取
dataset = dataset.shuffle(1000).batch(32).prefetch(tf.data.AUTOTUNE)

model.fit(dataset, epochs=10)

方案二:自定义Sequence扩展多文件

如果用自定义Sequence,可以先收集所有文件路径,再在__getitem__中分配每个batch对应的文件和行范围(需处理跨文件的batch拼接):

import pandas as pd
import numpy as np
import os
from tensorflow.keras.utils import Sequence

class MultiCSVBatchGenerator(Sequence):
    def __init__(self, data_dir, batch_size=32):
        self.batch_size = batch_size
        # 收集目录下所有CSV文件
        self.file_paths = [os.path.join(data_dir, f) for f in os.listdir(data_dir) if f.endswith(".csv")]
        # 统计每个文件的样本数(减去表头)
        self.file_sample_counts = []
        for path in self.file_paths:
            with open(path) as f:
                self.file_sample_counts.append(sum(1 for _ in f) - 1)
        self.total_samples = sum(self.file_sample_counts)
        self.total_batches = (self.total_samples + batch_size - 1) // batch_size

    def __len__(self):
        return self.total_batches

    def __getitem__(self, idx):
        start_sample = idx * self.batch_size
        end_sample = start_sample + self.batch_size
        x_batch = []
        y_batch = []
        current_sample = 0

        for path, cnt in zip(self.file_paths, self.file_sample_counts):
            if current_sample + cnt <= start_sample:
                current_sample += cnt
                continue
            # 计算当前文件需要读取的起始行和行数
            file_start = max(0, start_sample - current_sample) + 1
            file_end = min(cnt, end_sample - current_sample)
            read_rows = file_end - file_start + 1
            # 读取数据并拆分特征/标签
            df = pd.read_csv(path, skiprows=file_start, nrows=read_rows, header=None)
            x_batch.extend(df.iloc[:, :-1].values)
            y_batch.extend(df.iloc[:, -1].values)
            current_sample += cnt
            if len(x_batch) >= self.batch_size:
                break
        # 截断到指定batch size
        return np.array(x_batch[:self.batch_size]), np.array(y_batch[:self.batch_size])

# 初始化多文件生成器并训练
generator = MultiCSVBatchGenerator("./data_dir", batch_size=32)
model.fit(generator, epochs=10)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.05 06:05:31