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

如何在Keras在线分块训练中避免灾难性遗忘?

解决Keras分块训练灾难性遗忘的方案

核心问题分析

当前代码对每个数据块独立训练多轮(EPOCHS次),模型会过度拟合当前块的特征,完全覆盖之前从其他块学到的知识,这是导致灾难性遗忘的根本原因。


具体解决方法

1. 重构训练循环:全局多轮遍历所有数据块

不要在单个数据块上反复训练,而是将所有数据块视为一个大数据集的子批次,整体循环多轮:

TOTAL_EPOCHS = 10  # 原EPOCHS值,作为全局训练轮数
reached_training_targets = 0

for epoch in range(TOTAL_EPOCHS):
    # 每轮打乱数据块顺序,避免训练顺序偏差
    shuffled_fnames = np.random.permutation(input_file_names)
    for fname in shuffled_fnames:
        np_file = np.load(f"{TRAINING_FOLDER}/{fname}", mmap_mode='r')
        X = np_file['array1']
        y = np_file['array2']

        length_to_use = X.shape[0]
        reached_training_targets += X.shape[0]
        if reached_training_targets > NUM_SAMPLES:
            length_to_use -= (reached_training_targets - NUM_SAMPLES)
        if length_to_use <= 0:
            break

        X = X[:length_to_use]
        y = y[:length_to_use]

        rand_idx = np.random.permutation(X.shape[0])
        X = X[rand_idx]
        y = y[rand_idx]

        # 每个数据块仅训练1轮,全局完成TOTAL_EPOCHS轮遍历
        model.fit(X, y, epochs=1, batch_size=32, verbose=0, callbacks=[lr_schedule])
        np_file.close()
    print(f"全局训练轮次 {epoch+1}/{TOTAL_EPOCHS} 完成")

这种方式让模型逐步吸收所有数据块的信息,避免单块过度拟合。

2. 经验回放缓存:混合新旧数据训练

维护一个小型缓存池,每次训练新块时,混入部分旧块样本,强制模型保留旧知识:

replay_buffer = []
BUFFER_SIZE = 1000  # 根据内存调整缓存大小
reached_training_targets = 0

for fname in input_file_names:
    np_file = np.load(f"{TRAINING_FOLDER}/{fname}", mmap_mode='r')
    X = np_file['array1']
    y = np_file['array2']

    length_to_use = X.shape[0]
    reached_training_targets += X.shape[0]
    if reached_training_targets > NUM_SAMPLES:
        length_to_use -= (reached_training_targets - NUM_SAMPLES)
    if length_to_use <= 0:
        break

    X = X[:length_to_use]
    y = y[:length_to_use]
    rand_idx = np.random.permutation(X.shape[0])
    X, y = X[rand_idx], y[rand_idx]

    # 混合缓存中的旧数据
    if replay_buffer:
        replay_X, replay_y = zip(*replay_buffer)
        replay_X = np.concatenate(replay_X)
        replay_y = np.concatenate(replay_y)
        # 随机抽取部分旧样本
        sample_idx = np.random.permutation(len(replay_X))[:200]
        mix_X = np.concatenate([X, replay_X[sample_idx]])
        mix_y = np.concatenate([y, replay_y[sample_idx]])
        # 打乱混合后的数据
        mix_idx = np.random.permutation(len(mix_X))
        mix_X, mix_y = mix_X[mix_idx], mix_y[mix_idx]
    else:
        mix_X, mix_y = X, y

    # 训练混合数据
    model.fit(mix_X, mix_y, epochs=1, batch_size=32, verbose=2, callbacks=[lr_schedule])

    # 更新缓存,保留最新样本
    replay_buffer.extend(list(zip(X, y)))
    if len(replay_buffer) > BUFFER_SIZE:
        replay_buffer = replay_buffer[-BUFFER_SIZE:]
    
    np_file.close()

3. 增强模型正则化

通过正则化减少模型对单块数据的过度拟合:

  • 添加Dropout层:
    from tensorflow.keras.layers import Dropout
    
    # 在模型隐藏层后添加
    model.add(Dropout(0.2))
    
  • 添加L2正则化:
    from tensorflow.keras import regularizers
    
    model.add(Dense(64, activation='relu', kernel_regularizer=regularizers.l2(0.001)))
    

4. 优化学习率策略

采用预热+缓慢衰减的学习率,让模型逐步适应全量数据分布:

from tensorflow.keras.callbacks import LearningRateScheduler

def lr_scheduler(epoch, lr):
    if epoch < 5:
        return lr * 1.1  # 前5轮预热,逐步提升学习率
    else:
        return lr * 0.95  # 之后缓慢衰减

lr_schedule = LearningRateScheduler(lr_scheduler)

内容的提问来源于stack exchange,提问作者Valentino Zaffrani

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.29 22:30:22