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

如何为LSTM计算正确批量大小?避免时序数据丢失

解决LSTM时序数据批次划分时的剩余数据丢失问题

问题背景

我有1186条每日时间序列数据(包含CashIn、CashOut和Date字段),想要用LSTM模型预测2019-04-01至2019-04-30的CashIn和CashOut数值。我写了一个批量计算函数get_batches,尝试用序列长度30、批次大小40来划分数据集,但当把批次大小改为39时,会丢失最后16条数据,不想丢失这些数据,该怎么处理?

先贴出你当前使用的批次划分函数:

def get_batches(arr, batch_size, seq_length):
    batch_size_total = batch_size * seq_length
    n_batches = len(arr)//batch_size_total
    arr = arr[:n_batches * batch_size_total]
    arr = arr.reshape((batch_size, -1))
    for n in range(0, arr.shape[1], seq_length):
        x = arr[:, n:n+seq_length]
        y = np.zeros_like(x)
        try:
            y[:, :-1], y[:, -1] = x[:, 1:], arr[:, n+seq_length]
        except IndexError:
            y[:, :-1], y[:, -1] = x[:, 1:], arr[:, 0]
        yield x, y

核心问题分析

你的函数里len(arr)//batch_size_total是整数除法,直接截断了无法凑成完整批次的剩余数据,这就是换batch_size后丢数据的原因。下面给你几个实用的解决方案,既能保留所有数据,又能适配LSTM的训练需求:


解决方案1:修改批次逻辑,单独处理剩余数据

我们可以先处理完整的批次,再把剩余的数据单独生成小批次或者单个样本,确保每一条数据都被利用:

import numpy as np

def get_batches(arr, batch_size, seq_length):
    total_length = len(arr)
    batch_size_total = batch_size * seq_length
    n_full_batches = total_length // batch_size_total

    # 先处理所有完整批次
    for i in range(n_full_batches):
        start = i * batch_size_total
        end = start + batch_size_total
        arr_batch = arr[start:end].reshape((batch_size, -1))
        for n in range(0, arr_batch.shape[1], seq_length):
            x = arr_batch[:, n:n+seq_length]
            y = np.zeros_like(x)
            try:
                y[:, :-1], y[:, -1] = x[:, 1:], arr_batch[:, n+seq_length]
            except IndexError:
                y[:, :-1], y[:, -1] = x[:, 1:], arr_batch[:, 0]
            yield x, y

    # 处理剩余的零散数据
    remaining_start = n_full_batches * batch_size_total
    remaining_arr = arr[remaining_start:]
    
    # 确保剩余数据至少能生成一组(x,y)对(seq_length个输入+1个标签)
    if len(remaining_arr) >= seq_length + 1:
        # 尝试生成小批次
        small_batch_size = len(remaining_arr) // (seq_length + 1)
        if small_batch_size > 0:
            remaining_total = small_batch_size * (seq_length + 1)
            remaining_arr = remaining_arr[:remaining_total].reshape((small_batch_size, -1))
            for n in range(0, remaining_arr.shape[1] - seq_length):
                x = remaining_arr[:, n:n+seq_length]
                y = np.zeros_like(x)
                y[:, :-1], y[:, -1] = x[:, 1:], remaining_arr[:, n+seq_length]
                yield x, y
        else:
            # 剩余数据不够小批次,生成单个样本的批次
            for n in range(len(remaining_arr) - seq_length):
                x = remaining_arr[n:n+seq_length].reshape(1, -1)
                y = np.zeros_like(x)
                y[:, :-1], y[:, -1] = x[:, 1:], remaining_arr[n+seq_length].reshape(1, 1)
                yield x, y

解决方案2:用滑动窗口生成所有样本,灵活控制批次

如果不需要严格固定批次大小,推荐直接用滑动窗口生成所有可能的(x,y)样本,之后再用框架的数据集工具灵活设置批次。这种方法能100%利用数据,实现起来也更简单:

def get_sliding_window_samples(arr, seq_length):
    for i in range(len(arr) - seq_length):
        # 输入x是连续seq_length个数据
        x = arr[i:i+seq_length].reshape(1, seq_length)
        # 标签y是x的后移一位(对应LSTM的序列预测目标)
        y = arr[i+1:i+seq_length+1].reshape(1, seq_length)
        yield x, y

比如用TensorFlow的话,可以这样转换成可批量的数据集:

import tensorflow as tf

dataset = tf.data.Dataset.from_generator(
    lambda: get_sliding_window_samples(np.array(data_cashIn), 30),
    output_signature=(
        tf.TensorSpec(shape=(1, 30), dtype=tf.float32),
        tf.TensorSpec(shape=(1, 30), dtype=tf.float32)
    )
)
# 这里可以灵活设置批次大小,比如39或者40,不会丢数据
dataset = dataset.batch(39)

解决方案3:填充数据凑整(不推荐)

如果必须严格使用固定的batch_size和seq_length,可以对剩余数据进行填充(比如重复最后几个值、用均值填充),但这种方法会引入人工数据,可能影响模型泛化能力,仅在特殊场景下使用:

def pad_and_get_batches(arr, batch_size, seq_length):
    batch_size_total = batch_size * seq_length
    # 计算需要填充的长度
    pad_length = (batch_size_total - len(arr) % batch_size_total) % batch_size_total
    # 用数组最后一个值填充(也可以换成均值、中位数等)
    padded_arr = np.pad(arr, (0, pad_length), mode='edge')
    # 调用原来的函数处理填充后的数据
    return get_batches(padded_arr, batch_size, seq_length)

内容的提问来源于stack exchange,提问作者Emre BEGEN

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.13 07:58:04