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