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

如何基于Pandas CSV迭代器拆分输入输出迭代器适配Keras训练

大内存约束下Keras训练:分块CSV的输入输出拆分

问题背景

使用Keras训练神经网络时,因数据集规模过大(约425000行,每行含3600个输入特征+1个预期输出,共36001个元素),内存无法一次性加载全部数据,因此采用Pandas分块读取CSV:

df_iterator = pandas.read_csv("./formated_data/data-09-11-22.csv", chunksize=32, iterator=True)

需要将该迭代器拆分为输入数据迭代器和预期输出迭代器,喂给Keras的fit函数,且不能占用过多额外内存。自行实现了一个拆分函数,但不确定其正确性与内存友好性:

def split(iterable, n):
    iterators = []
    for i, iterator in enumerate(itertools.tee(iterable, n)):
        iterators.append(itertools.map(operator.itemgetter(i),iterator))
    return tuple(iterators)

你的拆分函数问题分析

你使用itertools.tee的实现会在内存中缓存迭代器元素——tee会把原迭代器产出的所有元素暂存起来,供多个派生迭代器读取。当处理大规模分块数据时,缓存的内容会持续占用内存,完全违背了分块读取的内存优化初衷,因此该方案不适合你的场景。

推荐方案1:自定义生成器(极简内存友好)

直接基于Pandas分块迭代器编写生成器,每次仅处理一个分块并拆分输入输出,无额外内存缓存:

import pandas as pd

def data_generator(df_iterator):
    for chunk in df_iterator:
        # 提取前3600列为输入特征,最后1列为预期输出
        X = chunk.iloc[:, :3600].values
        y = chunk.iloc[:, -1].values
        yield X, y

使用方式

# 注意:迭代器只能遍历一次,需重新初始化
df_iterator = pd.read_csv("./formated_data/data-09-11-22.csv", chunksize=32, iterator=True)
train_generator = data_generator(df_iterator)

# 需指定steps_per_epoch:总样本数 // 批次大小 = 425000 // 32 ≈ 13282
history = model.fit(train_generator, steps_per_epoch=13282, epochs=200, verbose=0)

推荐方案2:Keras Sequence类(支持验证集,更规范)

如果需要拆分训练集与验证集,Keras的Sequence类是更优选择——它是官方为分块数据设计的内存高效工具,支持数据打乱、验证集拆分、多进程加载:

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

class CSVDataSequence(Sequence):
    def __init__(self, csv_path, chunksize=32, is_train=True, validation_split=0.33):
        self.csv_path = csv_path
        self.chunksize = chunksize
        self.is_train = is_train
        self.validation_split = validation_split
        
        # 快速统计总样本数(不加载全量数据)
        self.total_samples = sum(1 for _ in open(csv_path)) - 1  # 减去表头行
        self.train_samples = int(self.total_samples * (1 - validation_split))
        self.val_samples = self.total_samples - self.train_samples
        
        # 预计算所有分块的起始行索引
        self.chunk_starts = list(range(0, self.total_samples, chunksize))
        
    def __len__(self):
        # 返回训练/验证的总步数
        if self.is_train:
            return (self.train_samples + self.chunksize - 1) // self.chunksize
        else:
            return (self.val_samples + self.chunksize - 1) // self.chunksize
        
    def __getitem__(self, idx):
        # 计算当前分块的起始行
        start_row = self.chunk_starts[idx]
        if not self.is_train:
            start_row += self.train_samples  # 验证集从训练集结束位置开始
            
        # 仅读取当前分块的数据
        chunk = pd.read_csv(
            self.csv_path,
            skiprows=start_row + 1,  # +1跳过表头
            nrows=self.chunksize,
            header=None  # 跳过表头后需指定无表头
        )
        
        X = chunk.iloc[:, :3600].values
        y = chunk.iloc[:, -1].values
        return X, y

使用方式

train_seq = CSVDataSequence("./formated_data/data-09-11-22.csv", is_train=True)
val_seq = CSVDataSequence("./formated_data/data-09-11-22.csv", is_train=False)

history = model.fit(
    train_seq,
    validation_data=val_seq,
    epochs=200,
    verbose=0
)

方案对比

  • 自定义生成器:代码简洁,适合无需验证集拆分的快速场景,内存占用极低,仅在当前分块加载时占用内存。
  • Sequence类:功能完整,支持验证集、数据打乱、多进程,适合正规训练流程,内存效率同样出色。

内容的提问来源于stack exchange,提问作者Bryan Carty

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.03 23:30:42