断点续训实战:边缘设备训练中断后如何无缝恢复
2026/9/7 22:53:49 网站建设 项目流程

最近我在嘉楠 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 权重保存的三种粒度与选择标准

在实际工程里,保存模型权重可以做成三种不同粒度:

  1. 完整训练检查点:包含模型权重、优化器状态、学习率调度器状态、epoch、batch 索引、随机数状态、验证集指标。这是续训的标准配置,体积最大,但恢复得最完整。
  2. 模型权重快照:只保存网络的 state_dict,不包含优化器。适合做迁移学习、模型融合、推理部署的前置检查。从快照续训不是不可以,但要手动调整学习率和优化器状态,风险高。
  3. 部署格式文件:比如 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(f"Epoch {epoch} Batch {batch_idx} Loss {loss.item():.6f}") return running_loss / len(train_loader) def main(): model = create_model(num_classes=10) optimizer = create_optimizer(model, lr=0.01) scheduler = create_scheduler(optimizer, step_size=10) train_loader = create_data_loader("data/train", batch_size=32) 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(f"Epoch {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_classes=10) optimizer = create_optimizer(model, lr=0.01) scheduler = create_scheduler(optimizer, step_size=10) train_loader = create_data_loader("data/train", batch_size=32) 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(f"Epoch {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.098732

loss 应该从和中断前差不多的量级继续下降,而不是大幅反弹。如果出现 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,通道数是 C,batch 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 开始收敛变慢,哪个阶段出现过中断和恢复,恢复后用了多久追回原来的指标。这些日志对后续调参和复现结果帮助很大。个人体会是,断点续训这套机制本身并不复杂,它考验的是对训练过程的细致程度。保存哪些字段、恢复哪些字段、保存多频繁、放到哪里,每一个决定都会影响后续训练的稳定性。建议你第一次做的时候,先在一个小数据集上跑通中断恢复全流程,确认一切正常之后,再放开手去跑正式训练,这样能省掉很多不必要的折腾。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询