Epoch

一个epoch指的是将整个训练数据集完整地遍历一次,完成一次向前计算及反向传播。由于完成一个epoch训练的周期较长(数据量大),一次性输入所有数据计算机无法负荷,所以将其分成多个batches。

Batch

Batch是每次输入网络进行训练的批次,batchsize是每个batch中训练样本的数量。

Iterations

iterations就是完成一次epoch所需的batch个数。

Epoch边界问题详解

在epoch的边界处,可能会出现多种需要特别注意的问题,这些问题统称为epoch边界问题。以下是具体的问题分类、原因分析及解决方案:

1. 数据分割与不完整批次处理

  • 问题描述
    当数据集大小无法被批次大小整除时,最后一个批次的数据量可能小于设定的批次大小(如总样本1000,批次大小32,最后一个批次仅8个样本)。
  • 潜在影响
    • 模型可能对小批次敏感,导致训练不稳定。
    • 统计指标(如损失、准确率)计算偏差。
  • 解决方案
    • 丢弃小批次:直接忽略最后的不完整批次,确保所有批次大小一致。
    • 填充数据:通过重复样本或补零填充最后一个批次至完整大小。
    • 动态调整批次大小:允许最后一个批次使用实际剩余样本数(需框架支持,如PyTorch的drop_last=False)。

2. 学习率调整与优化器状态

  • 问题描述
    学习率调度器(如余弦退火)通常在epoch结束时更新,若调整时机与epoch边界未对齐,可能导致学习率更新不准确。
  • 潜在影响
    • 学习率过早或过晚调整,影响模型收敛。
  • 解决方案
    • 严格按epoch调整:确保调度器仅在完整epoch后调用step()
    • 混合调度策略:结合批次级调整(如Warmup)和epoch级调整。

3. 模型保存与检查点管理

  • 问题描述
    在epoch结束时保存检查点时,未完整保存优化器、调度器状态,导致恢复训练失败。
  • 潜在影响
    • 训练中断后无法恢复,损失训练进度。
  • 解决方案
    • 保存完整状态:包括模型参数、优化器状态、学习率调度器状态、当前epoch和迭代步数。
    checkpoint = {
        'model': model.state_dict(),
        'optimizer': optimizer.state_dict(),
        'scheduler': scheduler.state_dict(),
        'epoch': epoch,
        'step': step
    }
    torch.save(checkpoint, 'checkpoint.pth')
    

4. 验证与评估模式切换

  • 问题描述
    在epoch结束后评估模型时,未正确切换模型为评估模式(如忘记调用model.eval()),导致BN层或Dropout层行为不一致。
  • 潜在影响
    • 验证集指标不准确(如BN层使用训练时的移动均值和方差)。
  • 解决方案
    • 显式切换模式:在验证前调用model.eval(),训练前调用model.train()
    # 训练循环
    model.train()
    for batch in train_loader:
        # 前向传播、反向传播...
    
    # 验证循环
    model.eval()
    with torch.no_grad():
        for batch in val_loader:
            # 前向传播...
    

5. 梯度累积的完整性

  • 问题描述
    使用梯度累积策略时,若在epoch结束时累积步数未达到预设值,可能导致梯度未完全应用。
  • 潜在影响
    • 参数更新不完整,影响模型收敛。
  • 解决方案
    • 强制更新:在epoch结束时,无论累积步数是否达标,均执行参数更新。
    • 调整累积步数:确保每个epoch的迭代次数是累积步数的整数倍。

6. 内存管理与资源清理

  • 问题描述
    每个epoch结束后未及时释放GPU缓存或临时变量,导致内存泄漏。
  • 潜在影响
    • 内存不足,训练意外终止。
  • 解决方案
    • 手动释放缓存:在epoch结束时调用torch.cuda.empty_cache()
    • 避免循环外引用:确保临时张量在作用域外被及时回收。

7. 分布式训练同步

  • 问题描述
    在分布式训练中,各进程未在epoch结束时同步,导致参数不一致。
  • 潜在影响
    • 模型参数在不同进程间出现分歧,训练效果下降。
  • 解决方案
    • 进程同步:使用torch.distributed.barrier()确保所有进程完成当前epoch。
    if config.distributed:
        torch.distributed.barrier()
    

总结

Epoch边界问题本质上是训练流程中状态管理与时序控制的挑战。解决这些问题需要:

  1. 明确生命周期管理:严格区分训练、验证、保存等阶段。
  2. 框架级支持:利用PyTorch等框架的特性(如DataLoaderdrop_last参数)。
  3. 代码健壮性:添加边界条件检查(如判断是否为最后一个批次)。
  4. 分布式协调:确保多设备/进程间的同步。
Logo

DAMO开发者矩阵,由阿里巴巴达摩院和中国互联网协会联合发起,致力于探讨最前沿的技术趋势与应用成果,搭建高质量的交流与分享平台,推动技术创新与产业应用链接,围绕“人工智能与新型计算”构建开放共享的开发者生态。

更多推荐