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

如何在低内存下从DataFrame生成滑动窗口训练集用于Keras训练?

问题描述

我有一个形状为(9177254, 7)的CSV文件,磁盘占用500MB。导入为数组时占用5GB内存,导入为DataFrame时约占用1GB内存。

训练集的每个样本应为形状为10000×7的数组,采用滑动窗口方式生成:第1-10000行是样本1,第2-10001行是样本2,以此类推,共生成9167254个样本。若生成完整的训练集数组(9167254×10000×7)将占用约10TB内存,这显然不现实。

当前使用的Keras训练代码如下:

model.compile(optimizer='adam',
              loss=tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True),
              metrics=['accuracy'])

history = model.fit(x_train, y_train, epochs=100, 
                    validation_data=(x_test, y_test))

我了解到训练输入x_test必须是数组,想知道是否有办法绕过这一限制?能否通过DataFrame的映射(或函数)来定义x_test以节省内存?


解决方案

完全不需要生成完整的数组,你可以通过自定义数据生成器或者TensorFlow Dataset API实现动态生成滑动窗口样本,全程无需预存所有样本,内存占用会极低。

方法1:使用TensorFlow Dataset API(推荐)

TensorFlow的tf.data.Dataset支持从DataFrame或磁盘文件直接构建数据流,动态生成滑动窗口,无需加载全量数据到内存。

步骤1:从DataFrame构建Dataset

假设你的DataFrame名为df,先转成TensorFlow Dataset:

import tensorflow as tf
import pandas as pd
import numpy as np

df = pd.read_csv("your_data.csv")
# 分离特征和标签(假设标签在最后一列)
features = df.iloc[:, :-1].values
labels = df.iloc[:, -1].values

feature_dataset = tf.data.Dataset.from_tensor_slices(features)
label_dataset = tf.data.Dataset.from_tensor_slices(labels)

步骤2:生成滑动窗口样本

利用window和flat_map生成固定大小的滑动窗口,过滤掉不完整的窗口:

window_size = 10000
# 生成特征窗口,shift=1表示每次滑动1行,drop_remain确保窗口长度为10000
windowed_features = feature_dataset.window(window_size, shift=1, drop_remainder=True)
windowed_features = windowed_features.flat_map(lambda window: window.batch(window_size))

# 标签对应每个窗口的最后一行,因此需要跳过前window_size-1个标签
windowed_labels = label_dataset.skip(window_size - 1)

# 合并特征和标签
dataset = tf.data.Dataset.zip((windowed_features, windowed_labels))

步骤3:划分训练/验证集并配置批处理

total_samples = len(df) - window_size + 1
train_size = int(0.8 * total_samples)

train_dataset = dataset.take(train_size)
val_dataset = dataset.skip(train_size)

# 设置批处理大小,开启预取提升训练效率
batch_size = 32
train_dataset = train_dataset.batch(batch_size).prefetch(tf.data.AUTOTUNE)
val_dataset = val_dataset.batch(batch_size).prefetch(tf.data.AUTOTUNE)

步骤4:用Dataset训练模型

直接将Dataset传入model.fit即可,无需转换为数组:

model.compile(optimizer='adam',
              loss=tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True),
              metrics=['accuracy'])

history = model.fit(train_dataset, epochs=100, validation_data=val_dataset)

方法2:自定义Keras Sequence生成器

如果你习惯用Keras的Sequence类,可以自定义生成器,每次只生成一个批次的样本,从DataFrame动态截取窗口:

from tensorflow.keras.utils import Sequence

class WindowDataGenerator(Sequence):
    def __init__(self, df, window_size, batch_size, label_col):
        self.data = df.values
        self.window_size = window_size
        self.batch_size = batch_size
        self.labels = df[label_col].values
        self.total_samples = len(df) - window_size + 1

    def __len__(self):
        # 返回总批次数量
        return self.total_samples // self.batch_size

    def __getitem__(self, idx):
        start_idx = idx * self.batch_size
        end_idx = start_idx + self.batch_size
        
        batch_x = []
        batch_y = []
        for i in range(start_idx, end_idx):
            # 截取当前窗口的特征
            window = self.data[i:i+self.window_size, :-1]
            batch_x.append(window)
            # 取窗口最后一行的标签
            batch_y.append(self.labels[i+self.window_size-1])
        
        return np.array(batch_x), np.array(batch_y)

使用生成器训练:

# 假设已划分好训练和验证DataFrame:df_train、df_val
train_generator = WindowDataGenerator(df_train, window_size=10000, batch_size=32, label_col='label')
val_generator = WindowDataGenerator(df_val, window_size=10000, batch_size=32, label_col='label')

model.fit(train_generator, epochs=100, validation_data=val_generator)

额外优化:直接从磁盘读取CSV

如果DataFrame加载后内存占用仍过高,可以直接用tf.data.experimental.make_csv_dataset从磁盘读取,全程不加载全量数据:

dataset = tf.data.experimental.make_csv_dataset(
    "your_data.csv",
    batch_size=1,  # 先按行读取
    label_name='label',  # 你的标签列名
    shuffle=False,
    num_epochs=1
)

# 转换为单元素流
dataset = dataset.unbatch()
# 后续步骤同方法1,生成滑动窗口、合并标签等

内容的提问来源于stack exchange,提问作者High GPA

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.23 14:45:46