断点续训实战:边缘设备训练中断后如何无缝恢复
最近我在嘉楠 AI Cube 上跑一个图像分类模型的训练数据量不算大但板上训练速度有限一个完整流程要跑十几个小时。前两次都是跑到七八个小时的时候因为 USB 供电不稳直接断连训练进程一停之前所有进度全部作废气得我差点把板子扔了。后来我把断点续训这套东西完整做了一遍核心就是加载已保存的模型权重在原有训练基础上继续迭代训练。跑通之后再也不用担心训练到一半断电、死机、内存溢出这些破事了。这篇就当是我自己踩坑后的一份总结同样在边缘设备上做训练的朋友可以直接抄作业。先说清楚它能解决什么问题。嘉楠 AI Cube 这类设备本质上是把 AI 训练和推理能力塞进一个非常小的嵌入式环境里性能比 PC 差一大截但胜在功耗低、便携、可以脱离云端的网络延迟独立干活。断点续训解决的就是在长时间训练过程中任何异常中断导致前面所有算力白费的问题。只要你在训练过程中定期把模型权重、优化器状态、当前迭代轮数这些东西落盘中断后就能从最近的存档点接着跑而不是从零开始。适合的场景包括数据集比较大、单次训练耗时很长、设备供电不稳定、需要反复调参试跑的人。下面我把整个思路、配置方法、代码实现和踩坑记录全部铺开来讲。1. 为什么在 AI Cube 上非做断点续训不可1.1 边缘设备训练的真实痛点很多人一提到训练模型默认就是 GPU 服务器、分布式集群。但嘉楠 AI Cube 这类设备主打的是端侧训练它用 RISC-V 核心加 KPU知识处理单元来加速神经网络计算。好处是成本低、无需高端显卡、可以在本地直接处理数据坏处也很直观——算力有限训练一个稍大的模型可能需要好几小时甚至一整天。训练时间一长各种意外就来了。最常见的是供电问题AI Cube 用 USB 供电电流稍微一波动板子直接重启其次是长时间跑训练导致内存碎片累积触发 OOM再有就是代码里没处理的异常比如某个 batch 读到了损坏的图片、日志目录写满、网络挂载目录超时。任何一次中断如果没有续训机制前面几个小时的迭代就全白跑了。我见过不少人想省事觉得大不了重新跑一遍。但训练不是线性过程后面的 epoch 是在前面所有 epoch 的基础上演化的你中断在第 15 个 epoch重新跑并不能保证跑到第 30 个 epoch 就能达到之前 15 轮再继续 15 轮的效果因为优化器状态、学习率退火曲线已经变了。所以从工程角度看断点续训不是可选项而是长时训练的基本配置。1.2 断点续训不只是保存一个权重文件很多新手第一次接触续训以为就是把 model 的权重存下来下次 load 一下就完事。实际上一个合格的断点续训机制至少要包含四样东西模型权重这是骨架承载了已经学到的特征。优化器状态包括动量、二阶矩估计等少了它训练会丢掉惯性和自适应调节能力。学习率调度器状态当前到了哪个 step、下一步该用多少学习率都要恢复。训练进度状态当前 epoch、当前 batch 索引、验证集最优精度、随机数生成器状态。打个比方你做饭做到一半不只是把锅里的菜装进保鲜盒你还得记住火候开到几档、盐放了半勺还是两勺、下一步是焖三分钟还是大火收汁。只存菜下次你再做就得靠猜。模型训练也是一样权重只是菜优化器状态才是火候和调料配比。很多人在续训时发现 loss 不但没降反而飙升很大概率就是优化器状态没恢复。随机数生成器状态这个很多人会忽略。训练中通常会做随机数据增强、随机打乱样本顺序如果每次续训都从相同的随机种子重新开始那么数据读取顺序会重复模型对某些样本的过拟合风险会增大。恢复随机状态后数据流能从断点平滑接续整个训练过程在概率意义上保持一致。1.3 方案选型为什么用 checkpoint 文件加续训脚本我对比过几种实现方案。第一种是干脆手动调低学习率用原来的权重初始化网络再从头跑这种方案最省事但学习率曲线不连续前期容易震荡中期收敛效率低第二种是模型并行备份每训练一步就同步把权重拷贝到多个位置这种方案开销太大在 AI Cube 这种资源紧张的设备上不现实第三种就是我最终采用的方案定期写 checkpoint 文件配套一个续训脚本启动时检测最新检查点恢复全部状态继续迭代。第三种方案的优势在于透明、可控、可移植。checkpoint 文件是一个独立的持久化实体放在 SD 卡或者计算机本地即使设备完全断电文件也不受影响。续训脚本可以独立运行不依赖训练脚本的交互式进程。另外这套机制不管是在 PC 上做原型验证还是迁移到其他设备上思路完全通用。2. 动手前必须搞懂的核心细节与参数2.1 嘉楠 AI Cube 的硬件分工CPU、KPU 与内存边界在写代码之前先了解一下 AI Cube 的硬件架构这样可以避免后面调参时两眼一抹黑。嘉楠 AI Cube 采用 K230 芯片平台CPU 是 RISC-V 双核异构设计有高性能大核和低功耗小核KPU 是专门做神经网络计算的单元支持卷积、池化、全连接这些常见算子。CPU 负责控制流程、做预处理和调度KPU 负责把计算密集的层拉走两边并行工作。内存方面AI Cube 用的是片内 SRAM 加外部 DRAM 的组合方式。训练过程中激活值、梯度、权重临时副本都住在内存里而 KPU 的本地存储只是加速计算的中间缓存。正因为内存总量有限batch size 不能拍脑袋设一个很大的值否则一个 step 下去直接内存爆满。这也是续训机制在设备上特别重要的原因之一内存越紧张进程就越容易在长时间运行后崩溃。理解了硬件边界你就知道 checkpoint 文件应该放在哪里最安全。建议放在可移动存储区比如 SD 卡而不是放在临时目录或内存文件系统里。设备重启后只有持久化存储里的文件还在。我见过有人把 checkpoint 写到 /tmp设备一重启文件全没了这个错误很低级但真的有人犯。2.2 权重保存的三种粒度与选择标准在实际工程里保存模型权重可以做成三种不同粒度完整训练检查点包含模型权重、优化器状态、学习率调度器状态、epoch、batch 索引、随机数状态、验证集指标。这是续训的标准配置体积最大但恢复得最完整。模型权重快照只保存网络的 state_dict不包含优化器。适合做迁移学习、模型融合、推理部署的前置检查。从快照续训不是不可以但要手动调整学习率和优化器状态风险高。部署格式文件比如 kmodel 或 ONNX这类格式主要给推理用做推理加速和端侧部署。它通常会做算子融合和量化不能当作训练检查点加载回来继续梯度更新。保存类型包含内容能否续训体积适用场景完整检查点权重优化器进度可以大长时间训练、断点恢复权重快照仅模型权重勉强中迁移学习、模型融合部署格式推理图量化参数不能小端侧推理部署在训练脚本里我一般每隔固定的 epoch 数就保存一份完整检查点同时额外导出一份权重快照。前者用于恢复后者用于随时评估和部署。两份文件都保留互不干扰。2.3 学习率与优化器状态怎么“接得上”续训时最容易翻车的就是学习率和优化器状态。很多框架在 new 一个 optimizer 对象时默认会把学习率重置为初始值。如果你只是加载了模型权重然后重新创建 optimizer结果就是用一个很大的学习率去继续一个已经收敛很久的模型loss 直接炸掉。解决思路有两个。第一在续训脚本里从 checkpoint 文件读取学习率调度器的当前值手动 set 到 optimizer 上。第二更稳妥的做法是加载 optimizer 的整体 state_dict因为 state_dict 里面除了会记录动量也会包含当前的学习率分组。加载之后可以用一段验证代码打出来看看确保学习的数值和中断前完全一致。还有一个我常用的技巧续训开始后的前几百个 step可以加一个小范围学习率 warmup比如从原来的三分之一线性升到目标值给模型一个缓冲期。因为即使你恢复了所有状态前一次的 batch 顺序和当前数据流之间还是可能有细微差异突然切换数据分布容易造成 loss 波动。warmup 可以让模型快速适应新流水线不会产生大的震荡。3. 完整实操在 AI Cube 上把续训跑起来3.1 环境准备与目录规划开始操作前先把环境捋清楚。我这里用的是基于 Python 的开发环境实际的接口名称以你手里的 SDK 版本为准但整体结构都差不多。需要准备的东西包括嘉楠 AI Cube 开发板刷好训练环境固件确保 CPU 和 KPU 驱动能正常加载。Python 环境里装好 numpy 和基础的科学计算库。训练数据集放在固定目录不要放在会被清理的临时目录。建立专门的 checkpoint 目录里面按时间戳或 epoch 存放检查点文件。我习惯的目录结构是这样的project/ ├── data/ │ ├── train/ │ └── val/ ├── ckpt/ │ ├── latest.ckpt │ ├── epoch_10.ckpt │ └── epoch_20.ckpt ├── train.py └── resume.py所有 checkpoint 文件名带轮数信息同时维护一个 latest.ckpt 符号链接指向最新文件。这样中断恢复时脚本只需要找 latest.ckpt 就行不用自己判断哪个文件是最新的。这个方法是我在 PC 上训练时留下的习惯放到 AI Cube 上一样好用。3.2 基础训练脚本关键代码为了把断点续训讲清楚这里给出一个近似的 Python 伪代码。实际运用时你只需要把 model、optimizer、scheduler 的接口替换成你正在用的框架对应名称。import os import json import random import numpy as np def save_checkpoint(state, filename): torch.save(state, filename) # 实际项目中替换成对应序列化接口 print(f[Checkpoint] saved to {filename}) def train_one_epoch(model, train_loader, optimizer, criterion, epoch): model.train() running_loss 0.0 for batch_idx, (data, target) in enumerate(train_loader): optimizer.zero_grad() output model(data) loss criterion(output, target) loss.backward() optimizer.step() running_loss loss.item() if batch_idx % 50 0: print(fEpoch {epoch} Batch {batch_idx} Loss {loss.item():.6f}) return running_loss / len(train_loader) def main(): model create_model(num_classes10) optimizer create_optimizer(model, lr0.01) scheduler create_scheduler(optimizer, step_size10) train_loader create_data_loader(data/train, batch_size32) criterion create_criterion() start_epoch 0 best_acc 0.0 for epoch in range(start_epoch, 30): train_loss train_one_epoch(model, train_loader, optimizer, criterion, epoch) val_acc evaluate(model, data/val) scheduler.step() print(fEpoch {epoch} done. Train Loss {train_loss:.6f} Val Acc {val_acc:.4f}) if val_acc best_acc: best_acc val_acc save_checkpoint({ model_state: model.state_dict(), optimizer_state: optimizer.state_dict(), scheduler_state: scheduler.state_dict(), epoch: epoch, best_acc: best_acc, rng_state: torch.get_rng_state(), }, ckpt/latest.ckpt)这个脚本的主循环里每个 epoch 结束都会做一次评估并在验证精度创新高时保存一次检查点。这里有个小细节我只在验证集精度刷新时保存而不是每个 epoch 都存。因为边缘设备存储空间有限频繁落盘不仅占内存还会拖慢训练。但如果你的设备存储充足建议至少每隔 5 个 epoch 强制保存一次以防验证精度长时间不刷新导致检查点老旧。3.3 断点续训脚本关键代码续训脚本和训练脚本的结构几乎一样差别在于启动阶段要做恢复操作。import os import torch def resume_checkpoint(model, optimizer, scheduler, ckpt_path): if not os.path.exists(ckpt_path): print([Resume] no checkpoint found, start from scratch) return 0, 0.0 ckpt torch.load(ckpt_path) model.load_state_dict(ckpt[model_state]) optimizer.load_state_dict(ckpt[optimizer_state]) if scheduler_state in ckpt: scheduler.load_state_dict(ckpt[scheduler_state]) start_epoch ckpt[epoch] 1 best_acc ckpt[best_acc] rng_state ckpt.get(rng_state) if rng_state is not None: torch.set_rng_state(rng_state) print(f[Resume] loaded checkpoint from epoch {ckpt[epoch]}, best acc {best_acc:.4f}) return start_epoch, best_acc def main(): model create_model(num_classes10) optimizer create_optimizer(model, lr0.01) scheduler create_scheduler(optimizer, step_size10) train_loader create_data_loader(data/train, batch_size32) criterion create_criterion() start_epoch, best_acc resume_checkpoint(model, optimizer, scheduler, ckpt/latest.ckpt) for epoch in range(start_epoch, 30): train_loss train_one_epoch(model, train_loader, optimizer, criterion, epoch) val_acc evaluate(model, data/val) scheduler.step() print(fEpoch {epoch} done. Train Loss {train_loss:.6f} Val Acc {val_acc:.4f}) if val_acc best_acc: best_acc val_acc save_checkpoint({ model_state: model.state_dict(), optimizer_state: optimizer.state_dict(), scheduler_state: scheduler.state_dict(), epoch: epoch, best_acc: best_acc, rng_state: torch.get_rng_state(), }, ckpt/latest.ckpt)注意 resume 函数里的三处关键恢复模型参数、优化器参数、调度器参数。如果检查点文件里没有调度器状态也不会报错但训练进度里学习率曲线就不连续了。我建议在保存检查点时就把这些字段写全一份标准格式的检查点既能给训练脚本用也能给后续的分析脚本用。3.4 训练中断后如何恢复操作步骤当设备中断后恢复到训练状态只需要三步。第一步把设备接回稳定电源确保供电没问题不要边充电边用不稳定的 USB 口。第二步检查 checkpoint 目录看看 latest.ckpt 是什么时候保存的如果距离中断时间比较久说明保存频率太低后面要把保存间隔调小。第三步运行续训脚本观察启动日志。我实际跑的时候第一次续训成功会看到类似这样的日志[Resume] loaded checkpoint from epoch 12, best acc 0.8423 Epoch 13 Batch 0 Loss 0.112356 Epoch 13 Batch 50 Loss 0.098732loss 应该从和中断前差不多的量级继续下降而不是大幅反弹。如果出现 loss 从零点几跳到三点几的情况那就要检查恢复逻辑了。启动之后建议先让它跑两三个 epoch确认稳定了再离开。不要一启动就丢下不管很多问题是在前几百个 step 内暴露的。3.5 验证续训是否成功的三条标准/#### 标准一损失值保持连续续训是否成功最直观的标准是 loss 曲线。中断前 loss 在 0.1 附近波动续训后第一个 step 应该也在 0.1 附近的量级最多因为数据流切换有小幅上升经过几个 batch 后回落到正常区间。如果 loss 一下子跳到几倍甚至几十倍说明权重没加载对或者学习率被重置。看完 loss 之后还要看它是否在接下来几个 epoch 里持续下降而不是原地抖动后者说明优化器状态没有正确恢复。标准二精度曲线持续上涨第二个标准是验证集精度。假设中断前最佳精度是 84.23%续训后大约两三个 epoch 内应该突破这个值至少也要接近。如果你发现精度明显低于断点时的数值比如掉到了 50%说明模型权重加载后出现了某种程度的参数错位。如果精度虽然不跌但好几个 epoch 一直不动大概率是学习率调度器状态丢了退 fire 到了一个极小的数值区间导致模型更新幅度太小。标准三日志中的迭代计数连续第三个标准是看 step 和 epoch 计数。中断前跑到 epoch 12续训后应该从 epoch 13 开始而不是从 0 开始。如果脚本忽略了这个计数虽然训练还能跑但学习率调度器会以为自己还在早期阶段用大的学习率去更新一个中后期的模型后果同样是 loss 飙高。我记得有一次就是忘了把 epoch 偏移量加进调度器看起来在用大学习率从头训实际上模型已经收敛得差不多折腾了大半天精度纹丝不动。4. 常见问题与排查技巧实录4.1 权重 key 不匹配加载直接报错这个问题特别常见。你在中断前用的模型结构是两层全连接加一个分类头结果中断后不知道谁改了代码重新定义模型时少加了一层那么 load_state_dict 就会抱怨 key 对不上。还有一种情况是分类类别数变了原本训练是 10 类现在模型定义成 20 类分类头的权重形状不一样直接加载失败。排查思路很简单打印出模型 state_dict 的 key 集合和 checkpoint 文件里的 key 集合两个集合做差集看多出来的或者少去的层是哪些。如果是分类头因为类别数变了可以选择不加载分类头的权重只加载 backbone 部分然后随机初始化新的分类头再继续训练。这个操作在迁移学习里叫 partial load但在续训场景下要谨慎因为分类头重新初始化意味着之前的类别区分能力全部清零通常只在类别定义确实变化时才使用。4.2 续训后 loss 猛涨比首训还高我遇到过一次续训启动后第一个 batch 的 loss 从 0.1 直接跳到 3.8当时我以为权重加载失败了于是把模型输出打印出来看发现输出值范围非常大典型的权重初始化被覆盖掉的现象。后来才发现续训脚本里我写错了保存顺序保存的是上一轮 epoch 之前的模型而不是上一轮 epoch 之后的模型相当于回退了一个 epoch 的状态。这个问题的本质是恢复的权重和优化器状态不一致。模型权重来自 epoch 30但优化器 state_dict 却是 epoch 29 的时候存的。两者本来就属于不同步的状态加载到一起自然会产生奇怪的梯度更新。排查的方法是保存时把 model、optimizer、scheduler 的状态放在同一个 dict 对象里原子写入读取时全部从同一个文件读出不要手动从两个文件拼状态。还有一种可能是保存频率太低检查点与中断点之间隔了太久数据分布已经发生变化。这种情况损失涨一点是正常的配合 warmup 一段时间就可以恢复。4.3 只保存权重导致优化器历史丢失有人说我只保存了模型权重续训的时候给 optimizer 重新初始化不行吗可以但代价很大。优化器里有两个东西是训练过程中逐步积累的动量项的累计梯度方向以及 Adam 自适应学习率里的二阶矩估计。这些信息不是初始化的零值它们是模型训练到当前状态的重要产物。以 Adam 为例如果二阶矩估计丢失优化器会用自己的初始值重建这会导致每个参数的学习率重新从默认值开始自适应。对已经收敛的模型来说这是一种很强的扰动。我记得有一次只加载权重loss 初始没有太大波动但训练了 5 个 epoch 后精度明显不如原来那次训练同阶段的成绩原因就是优化器状态丢失后自适应学习率走了完全不同的路径模型参数绕了一个大弯才回到正轨。所以我的结论是权重快照适合用来做评估和部署但真正要续训必须保存完整的 optimizer state。如果存储空间确实紧张我建议把优化器状态压缩后保存比如只保存 float16 版本恢复时再转回 float32精度损失很小但能省一半空间。4.4 AI Cube 续训中途又断掉的自动化处理一次续训成功并不代表万事大吉。训练继续进行的过程中依然可能再次遇到供电问题、内存占用增长、进程被系统杀掉。为了预防再次中断我在后面加了三层防护。第一层是更频繁地保存检查点我把保存间隔从每 10 个 epoch 改成每 2 个 epoch同时保留最新的 3 份文件循环覆盖旧文件。文件体积不大多占一点存储完全值得。第二层是加了一个简单的自动重启脚本。用 while 循环包裹整个训练进程一旦进程异常退出脚本会自动检测检查点文件然后重新运行续训脚本。这样即使半夜断电第二天早上发现训练可能已经自动恢复了。#!/bin/bash while true; do python train.py --resume ckpt/latest.ckpt echo training stopped, restarting in 5s... sleep 5 done第三层是在代码里捕获常见的 OOM 异常和系统信号。如果触发了内存泄漏导致的异常先把当轮状态保存一下然后再退出等外层脚本重启。这样做可以最大限度减少损失。4.5 内存不够batch size 到底怎么调在 AI Cube 上训练batch size 不是随便设置的。KPU 的算力主要用来做推理和反向传播中的矩阵计算但中间变量和梯度还是要放在内存里。当你发现训练跑到一半内存爆炸最先要调的就是 batch size。一个粗略的估算方式是这样的假设输入图片尺寸是 W x H通道数是 Cbatch size 是 B。那么单层卷积的输出激活值大小约为 B x C_out x H_out x W_out梯度也是同量级。多层的激活值累加起来再加权重梯度和优化器状态就是你单步训练的内存消耗。用这个公式反推先在电脑上跑一个小 batch统计内存占用再按比例缩放。我实际在 AI Cube 上用的 batch size 是 16 到 32。如果你用更大的模型或者更高分辨率的输入可能需要降到 8 甚至 4。调整 batch size 后最好同步调整学习率经验上 batch 减半学习率也减半这样收敛曲线比较稳定。如果内存还是不够还有一个办法是开启梯度累积。把 4 个 batch 的正向反向结果累加起来每隔 4 个 batch 做一次参数更新等效于用 batch size 64 在更新梯度但实际上在设备上只需要处理 batch size 16 的数据。这个方案对训练效果影响很小非常适合内存吃紧的嵌入式环境。这次在嘉楠 AI Cube 上把断点续训流程完整跑通之后我又顺手做了一件事把每份 checkpoint 文件里存的验证集精度汇总到一张表里训练完统一查看。这样就能看到模型从第几个 epoch 开始收敛变慢哪个阶段出现过中断和恢复恢复后用了多久追回原来的指标。这些日志对后续调参和复现结果帮助很大。个人体会是断点续训这套机制本身并不复杂它考验的是对训练过程的细致程度。保存哪些字段、恢复哪些字段、保存多频繁、放到哪里每一个决定都会影响后续训练的稳定性。建议你第一次做的时候先在一个小数据集上跑通中断恢复全流程确认一切正常之后再放开手去跑正式训练这样能省掉很多不必要的折腾。