STM32F103C8T6部署nano语言模型:64KB Flash定点推理
2026/9/18 19:14:43 网站建设 项目流程

1. 先把账算清楚:64KB Flash 和 20KB SRAM 到底能装多大的模型

在 stm32f103c8t6 上跑一个自训练的 nano 语言模型,这件事听起来像是个段子,但真做下来会发现,它难的从来不是"能不能跑",而是"跑多大的、跑多快、输出像不像人话"。我这次部署的目标很明确:一个字符级的小 Transformer,权重静态烧进 Flash,推理全程定点运算,串口吐字,整机不依赖任何外部存储。整个过程从训练到板端点亮大概花了我三个晚上,中间踩的坑比预想的多得多,尤其是量化和内存布局这两块。

如果你手上也有一块最小系统板,想搞明白大模型推理到底在算什么、定点部署到底难在哪,那这套流程能直接抄。它不适合拿来做产品,但作为理解 Transformer 推理链路的教具,性价比极高——你能亲眼看到每一个矩阵乘在 72MHz 的机器上花掉多少微秒。下面我把可行性核算、模型设计、权重导出、板端引擎、联调排查这几个环节完整拆一遍,stm32f103c8t6 的资源边界是贯穿全文的那条线。

1.1 参数量与 Flash 占用:先算账,再动手

很多人上手第一件事是去下载一个现成的 nanoGPT 然后往板子上烧,结果链接阶段就炸了。问题在于,STM32F103C8T6 的 64KB Flash 不是全给你放权重的,代码、启动文件、库函数、常量表都要占。我的经验值是:留给模型权重的空间不应该超过 40KB,剩下 24KB 要留给代码和可能的升级余量。

先建立一个参数量估算公式。字符级 Transformer 的参数主要来自四块:

  • 词嵌入矩阵:vocab_size × d_model(如果做权重共享,输出头不额外占空间)
  • 位置嵌入:ctx_len × d_model
  • 每个 Block 的注意力投影:4 × d_model²(Q、K、V、O 四个矩阵)
  • 每个 Block 的前馈层:2 × d_model × d_ff

以我最终采用的配置为例:词表 48、d_model=48、2 层、d_ff=96、上下文长度 24。逐项加起来是 2304(词嵌入)+ 1152(位置嵌入)+ 2×9216(注意力四矩阵)+ 2×9216(前馈两层)+ 若干归一化的缩放参数,总参数量落在40.7K 左右,int8 量化后大约 40KB 权重。这个数字是卡着上限设计的。

为了让你有个直观对比,我把试过的三套配置列出来:

配置代号d_model层数d_ff上下文参数量int8 权重结论
入门版3226416约 19K约 19KB轻松放下,但表达能力很弱
标准版4829624约 40.7K约 40KB卡上限,需精打细算
激进版64319232约 130K约 127KB64KB 根本放不下

那张"激进版"的表我建议你留着,以后想换芯片的时候一眼就能看出差多少。参数量对层数和 d_model 是平方级敏感,对上下文长度是线性敏感,所以砍 d_model 是最快的瘦身手段,砍上下文长度则是性价比最低的——上下文太短,模型连"上一句话"都记不住,生成出来就是胡言乱语。

1.2 SRAM 的真正大头是 KV Cache 和中间激活

Flash 只是第一道坎,20KB SRAM 才是真正让人睡不着的部分。这里必须区分两种内存用法:权重是只读的,可以常驻 Flash;激活值必须放 RAM,而且每生成一个 token 都要重算。

推理时的 RAM 占用主要来自五块:

  • KV Cache:2 × 层数 × 上下文长度 × d_model字节。标准版配置下是2×2×24×48 = 4608字节,int8 存储。
  • 隐藏状态与残差流:用 int16 保存残差,24 × 48 × 2的话开销不小,实际我只保留当前 token 的一行,48 × 2 = 96字节。
  • 前馈中间层:d_ff个激活,96 × 2 = 192字节。
  • 注意力分数缓冲区:头数 × 上下文长度 × 4字节,2×24×4 = 192字节。
  • 输出 logits:48 × 4 = 192字节(int32 累加器)。

再加上一层 GEMM 的输出暂存、串口 DMA 缓冲区和栈空间,实测静态占用大约 9KB,栈我给到 2KB,加起来 11KB 左右。这个余量是刻意留的,因为调试阶段你会忍不住加日志、加打点、加缓冲区,20KB 很容易就爆。我踩过一次栈溢出,表现是随机 HardFault,查了两小时才发现是局部数组开太大——启动文件里Stack_Size默认只有 0x400(1KB),必须改。下面这段是启动文件里的关键改动:

Stack_Size EQU 0x00000800 ; 从 1KB 提到 2KB Heap_Size EQU 0x00000200 ; 堆基本不用,留一点点

注意:栈溢出的症状非常"随机",可能改一行无关代码就正常了,也可能是同一个函数偶尔崩。凡是遇到无规律的 HardFault,先把栈加大一倍试试,能省下大量时间。

1.3 Cortex-M3 没有 DSP 双乘加,算力上限要按实际情况估

这是最容易被误导的一点。网上大量"在单片机上跑神经网络"的教程用的是 STM32F4 或 GD32F303,这些是 Cortex-M4 内核,带 DSP 扩展指令,SMLAD一条指令能在一个周期里做两个 16 位乘加。STM32F103C8T6 是 Cortex-M3 内核,没有 DSP 扩展,没有 SMLAD,也没有硬件浮点单元。它有的是SMLAxy这类单路 16 位乘加,以及硬件SDIV/UDIV。M3 同样没有 SIMD,所以任何"双通道 MAC"的优化思路在这里都不成立。

那实际吞吐怎么估?标准版配置下,每个 token 的乘加次数大约是:

  • 注意力四个投影矩阵:4 × 48 × 48 = 9216
  • 注意力分数计算:头数 × 上下文 × 头维度 = 2 × 24 × 24 = 1152
  • 前馈第一层:48 × 96 = 4608
  • 前馈第二层:96 × 48 = 4608
  • 输出投影(与词嵌入共享):48 × 48 = 2304

合计约22K 次乘加。听起来不多,但 M3 上一条 int8 乘法在 C 里的实际开销是:LDRSB取数、SMULBB相乘、ADD累加,再加上循环计数,乐观估计3 到 5 个周期一次 MAC。按 4 周期算,22K × 4 = 88K 周期,72MHz 下就是1.2ms。这是纯理想值,实际还有归一化、指数运算、量化重定标、循环控制、Flash 访问等待周期等开销,我实测的稳态速度落在20 到 60ms 每 token(编译器-O2且循环写好之后)。生成 20 个字符大约 0.5 到 1.2 秒,这个速度用于演示完全够用。

顺便说一句,Flash 等待周期是个隐形杀手。72MHz 主频下必须配 2 个等待周期,并且打开预取缓冲,否则不仅慢,还可能直接跑飞。这段配置在system_stm32f1xx.c里,CubeMX 生成的代码通常会处理,但你如果手工改过时钟树,一定要回头确认:

/* Flash 延迟配置:72MHz 必须 2 等待周期 + 预取使能 */ FLASH->ACR |= FLASH_ACR_LATENCY_2 | FLASH_ACR_PRFTBE;

还有个小技巧:把最内层 GEMM 循环用__attribute__((section(".RamFunc")))或 Keil 的__ramfunc放到 RAM 里执行。F1 系列是哈佛结构,指令走 I-Code 总线、数据走 D-Code 总线,把取指挪到 RAM 的 System 总线上,就不再和从 Flash 取权重抢带宽。我改完之后,单 token 耗时从 78ms 降到了 51ms,这个收益很实在。


2. 模型设计:在 PC 上训练一个"故意做小"的字符级 Transformer

模型这块的心态要摆正:你不是在追求效果,你是在追求"刚好能被 64KB 装下、还能吐出人能看懂的东西"。这个约束反过来决定了很多设计选择,其中最关键的三个决策是字符级词表、权重共享、以及刻意压小的上下文长度。

2.1 词表与语料:字符级 + 小语料是唯一划算的路线

BPE 分词在小模型上是奢侈品。假设你用 512 的词表,光词嵌入矩阵就是512 × 48 = 24576个参数,占总参数预算的六成,而它带来的压缩收益在这个规模下根本体现不出来。字符级词表就不一样了,我的语料里实际只出现了 41 个不同字符,加上几个预留符号,词表定成48,词嵌入只占 2304 个参数。

语料的选择也有讲究。我用的是自己攒的一批结构化短文本:几千条固定格式的记录行,字符集收敛、句式重复度高。这种语料的好处有两个:一是词汇表能压到 48 以内,二是模型很容易学到语法骨架,即使参数量只有 4 万,生成出来的东西也是有模有样的格式串而不是乱码。如果你拿《莎士比亚全集》去训,40K 参数的模型学到的只会是字母频率分布,输出一堆像英文但不成词的垃圾。

语料规模我控制在200KB 到 500KB 纯文本,也就是二十万到五十万字符。这个量级在 CPU 上训练十几分钟就能收敛,用显卡更快。

提示:一定要先跑一次字符频率统计,把低频字符(出现次数少于 50 次的)统一映射到一个<unk>标记,否则词表会被一堆只出现一两次的符号撑大,白白浪费参数预算。

2.2 结构定稿:2 层 48 维、24 上下文,附参数量明细

网络结构我最后定成了这样:2 层 Transformer Block,d_model=48,2 个注意力头(头维度 24),前馈隐藏层 96,上下文窗口 24,pre-norm 结构(LayerNorm 放在子层前面)。Pre-norm 在浅层小模型上比 post-norm 稳得多,训练不容易发散,这个选择在只有两层的时候差别很明显。

最后一层的输出投影和词嵌入共享权重,这是省参数的常规操作,能省下 2304 个参数;在 4 万参数的总量下,这已经是 5% 的节省,值得做。

具体到每个层级的张量形状,我整理成表方便你对照写导出脚本:

张量名形状参数量说明
token_emb48 × 482304与输出头共享
pos_emb24 × 481152可学习位置编码
attn.qkv48 × 1446912把 Q、K、V 合并成一个矩阵,减少访存
attn.out48 × 482304输出投影
ffn.fc148 × 964608升维
ffn.fc296 × 484608降维
norm 缩放/偏置各 48约 384两层各一组

把 Q、K、V 合并成一个 48×144 的大矩阵是个实用技巧:一次 GEMM 就能算出三个结果,激活值只需读一遍,在内存带宽吃紧的 M3 上,这比三次独立 GEMM 快不少。

2.3 训练超参与"故意过拟合"策略

训练阶段我用的是一套比较激进的超参,目的就是让小模型在小语料上快速收敛到"能背下来一部分"的程度:

参数取值理由
优化器AdamW小模型上比 SGD 收敛快得多
学习率3e-3,余弦退火到 3e-4小模型能吃大学习率
权重衰减0.01抑制过拟合,但别太大
批大小6424 长度序列,显存/内存压力小
训练步数15000 到 30000看损失曲线收敛情况
Dropout0.0 到 0.1语料小的时候直接开 0
梯度裁剪1.0防偶发梯度爆炸

有个反直觉的点:语料只有几百 KB 的时候,我建议直接把 Dropout 设成 0。因为此时你的目标不是泛化,而是让模型把语料的分布背得足够扎实。40K 参数的模型在 300KB 文本上根本记不住全部内容,过拟合风险很低,反倒是欠拟合更容易出现。我第一版老老实实开了 0.1 的 Dropout,结果损失卡在 2.4 下不去,改成 0 之后训练损失直接掉到 1.1,生成质量肉眼可见地变好。

训练时还要盯一个指标:验证集的字符级困惑度。如果训练损失一直在降但生成质量没变好,说明模型开始死记硬背了。不过对于这个规模的演示项目,我一般只看训练损失加人工目测生成样本,够用。

2.4 训练阶段就要做量化模拟,别等到板子上才发现对不上

这一步是我最想强调的。很多人训练完直接导出 int8 权重,烧进板子,然后发现输出是乱码,再回头一层层查,最后发现是量化误差累积或者有个算子实现错了。正确做法是:在 PyTorch 里先做一次"假装量化"的前向,把量化误差的影响提前暴露出来。

具体做法是给每个权重张量加一个量化-反量化的包装:

def fake_quant_per_channel(w, n_bits=8): # w: [out_ch, in_ch],按输出通道做对称量化 qmax = 2 ** (n_bits - 1) - 1 # 127 scale = w.abs().amax(dim=1, keepdim=True) / qmax scale = scale.clamp(min=1e-8) q = torch.round(w / scale).clamp(-qmax, qmax) return q * scale, scale

把这个包装套到所有线性层和前馈层上,用同样的权重跑一遍前向,对比原始输出的 top-1 token 一致率。如果一致率在 95% 以上,说明量化方案没问题,可以放心导出;如果掉到 80% 以下,说明你的量化粒度太粗,需要改成 per-channel 或者提高激活值的位宽。我当时第一次用 per-tensor 量化,一致率只有 72%,换成 per-channel 之后直接到 97%,生成质量天差地别。

注意:PyTorch 的torch.round是"四舍五入到偶数"(banker's rounding),而 C 里如果直接用(int)(x + 0.5f)是"四舍五入"。这两种行为在 0.5 的边界上会给出不同结果。虽然实际影响很小,但如果你在排查"PC 和板子输出不一致",这是个值得检查的点。我建议导出时统一用np.floor(x + 0.5),然后 C 侧也用同样的规则。


3. 权重导出:从 state_dict 到 C 数组的完整流程

训练完之后,PyTorch 的权重是 float32 的,需要压成 int8 并且变成 C 能直接用的数组。这一步看着简单,实际上坑最多,因为它同时牵扯量化精度、内存布局、以及链接器怎么摆放这些数据。

3.1 量化方案:per-channel 对称 int8,激活值动态定标

权重量化我用的是对称 per-channel int8,也就是对每个输出通道单独算一个缩放因子。公式很简单:

scale[w] = max(|W[w, :]|) / 127 q[w, i] = round(W[w, i] / scale[w]),并裁剪到 [-127, 127]

用 -127 到 127 而不是 -128 到 127,是为了保持对称性,让反量化的时候没有额外的零点偏移,能省一条指令。这个细节在 M3 上值得在意。

激活值的量化我用动态 per-tensor:每一层的输入在运行时求一次最大值,算出缩放因子,再把 int8 激活值重定标到下一层需要的尺度上。这样做精度比静态定标好,代价是每层多一次求最大值的遍历——对 48 维向量来说,这点开销可以忽略。

重定标的实现要小心溢出。累加结果是 int32,等比缩放系数我用 Q15 定点数保存,乘法前必须提升到 64 位:

/* acc: int32 累加结果;ratio_q15: Q15 缩放比;返回 int8 量化值 */ static inline int8_t requant_i8(int32_t acc, int32_t ratio_q15) { int64_t t = (int64_t)acc * ratio_q15; /* 最大约 1e6 * 32767 ≈ 3.3e10,必须 64 位 */ t = (t + (1 << 14)) >> 15; /* 带四舍五入的右移 */ if (t > 127) t = 127; if (t < -127) t = -127; return (int8_t)t; }

这里的ratio_q15 = (scale_w × scale_x) / scale_out,在每一层推理前算一次,不用每次都算。M3 上一条 64 位乘法的开销是几个周期,但每层只用几十次,完全不是瓶颈。

3.2 导出脚本:生成 model_weights.h 的完整逻辑

导出脚本我写成了单个 Python 文件,核心逻辑是遍历state_dict,对每个张量做量化、转成 int8 或者 int16 数组、生成 C 头文件。结构大概是这样:

import numpy as np def emit_array(f, name, arr, dtype="int8_t"): flat = arr.flatten() f.write(f"static const {dtype} {name}[{flat.size}] = {{\n") for i in range(0, flat.size, 16): chunk = ",".join(str(int(v)) for v in flat[i:i+16]) f.write(" " + chunk + ",\n") f.write("};\n\n") with open("model_weights.h", "w") as f: for name, w in quantized_weights.items(): emit_array(f, name, w) # 缩放因子单独用 float 存,只读不参与运算,不占 RAM for name, s in scales.items(): emit_array(f, name + "_scale", s, dtype="float")

生成的数组全部声明为static const,这一点极其重要,下一节展开讲。

顺便说下体积。40K 个 int8 权重写成十进制文本,源文件大概 160KB 到 200KB,编译器处理起来毫无压力。如果你嫌源文件太大影响编辑器打开速度,可以改成每行 32 个数,或者把权重拆成多个头文件按层分开。

3.3 const 关键字与链接脚本:Flash 里放和 RAM 里放的区别

这是我在这个项目上踩的最大的一个坑,值得单独拿出来说。

在 C 语言里,const修饰的全局数组会被放进只读段,链接器把它安排在 Flash 地址空间,不占用 RAM。但如果你漏了const,比如写成static int8_t w_attn[6912] = {...},它就会被归到.data段——这个段的初始化数据同样存在 Flash 里,但启动时会被__main里的散列加载代码整个复制到 RAM。结果就是:RAM 瞬间被吃掉 40KB,而芯片只有 20KB,链接阶段直接报 "region RAM overflowed"。更隐蔽的情况是数组比较小,链接过了,但运行起来各种诡异行为,因为栈被挤到了别的区域。

所以检查清单里第一条永远是:所有静态权重数组必须带const

第二件事是确认链接脚本和 map 文件。Keil 里看.map文件,或者用fromelf --text -z看段分布;GCC 工具链下用arm-none-eabi-size-Wl,-Map=out.map。你要确认的数字是:.text+.rodata+.constdata加起来不超过 64KB,.data+.bss加栈不超过 20KB。我当时的分布是这样的:

大小说明
.text(代码)约 14KB含 HAL 库、串口、主逻辑
.rodata(权重)约 40KBint8 权重 + 缩放因子表
.data约 200B少量带初值的全局变量
.bss约 9KBKV Cache、激活缓冲区
2KB启动文件里配置

总共 62KB 左右,卡得死死的。所以你如果按标准版配置来,务必先把 HAL 库里不用的模块关掉,特别是HAL_ADCHAL_SPIHAL_TIM这些,CubeMX 里不勾就不会编译进去。

提示:GCC 工具链下建议加上-ffunction-sections -fdata-sections-Wl,--gc-sections,能把没引用到的函数和数据整段删掉,我这边省了差不多 6KB。Keil 的 ARM Compiler 6 对应选项是-ffunction-sections -fdata-sections加链接器的--remove

3.4 实在装不下:int4、剪枝和换芯片的三条路

如果你的模型怎么压都装不进去,有三条路可以走,按投入产出比排序:

第一条是int4 权重量化。只对前馈层和注意力投影做 4 位量化,词嵌入和归一化参数保持 8 位。体积直接减半,40KB 变 20KB。代价是推理时要多一步解包,把两个 4 位值从一个字节里拆出来,每字节多花 1 到 2 个周期,整体速度大概慢 15%。精度损失方面,在 4 万参数这个量级上,per-channel 的 int4 效果并没有想象中那么糟,我实测困惑度只涨了 0.08。这个方案适合 Flash 差一点点的情况。

第二条是结构化剪枝。把前馈隐藏层从 96 砍到 64,或者干脆砍掉一层。这是最直接的瘦身方式,但会实质性地影响模型能力,属于"改模型"而不是"改表示",要重新训练。

第三条是换芯片。这里有个便捷的选择:pin-to-pin 兼容的国产替代芯片里,有不少型号在同样的 LQFP48 封装下提供了更大的 Flash 和 RAM,有的还升级到了 Cortex-M4 内核带 DSP 指令。换上去之后代码基本不用改,还能用上SMLAD这类双乘加指令,速度能有明显提升。如果你本来就打算把项目做下去,这条路比死磕 64KB 划算得多。当然,如果你就是要挑战极限,那 64KB 里塞下 4 万参数模型已经很有成就感了。


4. 板端推理引擎:从 GEMM 到串口吐字

板端的代码我大概写了八百行,其中真正核心的就是一个 int8 矩阵乘、一个归一化、一个近似指数、以及一层循环调度。这部分我把关键实现和踩坑点都列出来,你照着写基本能复现。

4.1 工程骨架与 CubeMX 关键配置

工程我用 CubeMX 生成,编译器是 Keil MDK(ARM Compiler 6),也试过 STM32CubeIDE 的 GCC,两边都可以,ARMCC 生成的代码在 GEMM 内循环上大概快 10%。配置上有几个必改项:

  • 时钟:HSE 8MHz 外部晶振,PLL ×9,SYSCLK 72MHz,APB1 36MHz,APB2 72MHz。
  • Flash:2 个等待周期 + 预取使能。
  • USART1:PA9/PA10,115200 波特率,开启 DMA 发送,避免串口阻塞打断推理节奏。
  • SysTick:保持 1ms 中断,用于粗略计时;精确计时用 DWT 的 CYCCNT。
  • 调试口:SWD,保留,方便看变量。注意别在正式版里关掉 SWD,否则调试很不方便。

DWT 计时器的初始化只有三行,但非常值得加:

CoreDebug->DEMCR |= CoreDebug_DEMCR_TRCENA_Msk; DWT->CYCCNT = 0; DWT->CTRL |= DWT_CTRL_CYCCNTENA_Msk; /* 用法:uint32_t t0 = DWT->CYCCNT; ... uint32_t dt = DWT->CYCCNT - t0; */

有了它,你能精确定位到是哪一个算子慢,比用毫秒级计时器靠谱得多。

4.2 int8 GEMM 的写法与寄存器分块

矩阵乘是整个推理的绝对热点,这部分值得反复打磨。最朴素的版本三层循环,一行一个输出:

void gemm_i8_naive(const int8_t *W, const int8_t *x, int32_t *out, int out_ch, int in_ch) { for (int o = 0; o < out_ch; o++) { const int8_t *w = W + (size_t)o * in_ch; int32_t acc = 0; for (int i = 0; i < in_ch; i++) { acc += (int32_t)w[i] * (int32_t)x[i]; } out[o] = acc; } }

这个版本的问题是每一行都要把输入向量x重新读一遍。x只有 48 字节,能全部待在寄存器里吗?M3 只有 16 个通用寄存器,放不下。所以更好的做法是一次算 4 行输出,共享对x的读取,这就是所谓的寄存器分块:

void gemm_i8_4row(const int8_t *W, const int8_t *x, int32_t *out, int out_ch, int in_ch) { int o = 0; for (; o + 4 <= out_ch; o += 4) { const int8_t *w0 = W + (size_t)(o + 0) * in_ch; const int8_t *w1 = W + (size_t)(o + 1) * in_ch; const int8_t *w2 = W + (size_t)(o + 2) * in_ch; const int8_t *w3 = W + (size_t)(o + 3) * in_ch; int32_t a0 = 0, a1 = 0, a2 = 0, a3 = 0; for (int i = 0; i < in_ch; i++) { int32_t xv = x[i]; a0 += w0[i] * xv; a1 += w1[i] * xv; a2 += w2[i] * xv; a3 += w3[i] * xv; } out[o + 0] = a0; out[o + 1] = a1; out[o + 2] = a2; out[o + 3] = a3; } /* 处理不足 4 行的尾巴 */ for (; o < out_ch; o++) { /* 单行版本 */ } }

实测这个改动能把 GEMM 部分的速度提升接近 40%,因为x[i]的加载次数降到了四分之一,而 M3 上LDRSB是实打实要花周期的。

还有两点值得注意:内层循环要顺序访问权重。F1 的 Flash 预取缓冲是按顺序优化的,如果你在内层跳着取数,预取基本失效,每个字节都要吃满等待周期。所以权重一定要按行主序排列,内层循环遍历连续的in_ch维度。另外,int32_t累加器的溢出要算一下:127 × 127 × 48 ≈ 774K,远小于 int32 上限,安全。

4.3 其余算子:归一化、GELU 近似、Softmax 与采样

GEMM 之外还有几个算子需要处理,每一个都有陷阱。

LayerNorm(我这里用 RMSNorm 简化版):需要算均方根,涉及开方。M3 有硬件除法但没有硬件开方,用牛顿迭代两次就够精度了。为了省事也可以直接查表,把均方根值量化到 8 位索引,用 256 项的倒数表,这样连除法都省了。我用的是迭代法,实测每个 token 花不到 100 个周期。

GELU 激活函数:原版 GELU 带erf,在定点上算很痛苦。我直接用x * sigmoid(1.702x)这个近似,再退一步用硬 sigmoid(clamp(x/6 + 0.5, 0, 1))配合 Q15 定点乘,效果在这个规模下肉眼看不出来差别。在小模型上,激活函数的精度不是瓶颈,量化误差才是。

Softmax:这是唯一让我纠结的地方。需要指数运算,但没有 FPU,也不能引入expf——那会把整个数学库拖进来,代码体积暴涨十几 KB。我的做法是:先减去最大值(数值稳定性必须做,否则 int32 累加后再指数会溢出),然后用一张 128 项的定点指数表加线性插值。表放在 Flash 里,只占 256 字节。

采样:温度系数和 top-k 都是小整数运算。随机数我用 xorshift32,种子取自 SysTick 的当前值加一个固定常量,保证每次上电输出不一样但又可复现(想复现就固定种子):

static uint32_t rng_state = 0x12345678u; static inline uint32_t xorshift32(void) { uint32_t x = rng_state; x ^= x << 13; x ^= x >> 17; x ^= x << 5; rng_state = x; return x; }

温度的实现方式是把 logits 除以温度对应的定点系数。温度 0.8 大概对应放大 1.25 倍。top-k 我用的是简单版本:扫一遍找第 k 大的值作为阈值,把低于阈值的 logits 压到极低,再按概率采样。k=8 的时候这个扫描的开销可以忽略。

4.4 主循环、串口输出与耗时统计

整体推理循环的结构很清晰:把 prompt 逐个 token 喂进去建立 KV Cache,然后进入自回归生成,每生成一个 token 就通过串口发出去。

/* 1. 预填充:把 prompt 的每个字符过一遍前向,填充 KV Cache */ for (int i = 0; i < prompt_len; i++) { int tok = char_to_id(prompt[i]); forward(tok, i, /*write_kv=*/1); } /* 2. 自回归生成 */ for (int step = 0; step < max_new_tokens; step++) { int pos = prompt_len + step; int tok = last_token; const int32_t *logits = forward(tok, pos, 1); int next = sample_topk(logits, TEMP_Q8, TOP_K); uart_send_byte((uint8_t)id_to_char(next)); last_token = next; }

我把耗时统计做成了编译期开关,发布版关掉,调试版打开,每次生成完在串口打个汇总:

#ifdef PROFILE_ENABLE uint32_t cyc = DWT->CYCCNT - t_start; /* 72MHz 下,每个周期约 13.9ns */ printf("[perf] tokens=%d cycles=%lu avg_us=%lu\r\n", n, cyc, (unsigned long)(cyc / 72)); #endif

注意:printf如果支持%f会引入浮点格式化代码,这部分在 ARMCC 里能占好几 KB。要么用整数格式化,要么在 Keil 的 Target 选项里勾选 "Use MicroLIB",体积能小一截。

4.5 实测性能与调优空间

我把几个版本的数据整理成了表,方便你判断自己的实现处于什么水平:

优化阶段单 token 耗时说明
初版 C 代码,-O0约 310ms完全不能用,调试阶段
-O2 编译约 130ms编译器优化贡献最大
4 行分块 GEMM约 78ms复用输入向量
GEMM 搬到 RAM 执行约 51ms减少 Flash 取指竞争
归一化改定点 + 指数查表约 34ms去掉软浮点和除法
激活改 int16 残差通路约 28ms减少重定标次数

最终稳定在 28 到 35ms 每 token,生成 20 个字符不到 0.7 秒。这个速度在串口上看起来像是"打字机"效果,挺有意思。

还能再快吗?理论上还有空间:内层循环用汇编写、把 KV Cache 也搬进 RAM、用多字节一次读取的技巧。但我实测下来收益递减,而且代码可读性急剧下降。28ms 已经足够让人接受,再压下去性价比不高。


5. 对拍与排查:PC 输出对不上的那些坑

这部分是全文我觉得最有价值的地方。定点推理最大的痛苦不是"写不对",而是"不知道自己哪里不对"——输出是一堆乱码,但不报错、不崩溃,你没有任何线索。

5.1 对拍方法:同一套 C 代码先在 PC 上跑

我的方法很土但极其有效:把板端那套推理引擎(不含硬件相关部分)原封不动地编译到 PC 上运行。C 语言的可移植性在这里帮了大忙,只要把int8_tint32_t这些类型定义好,去掉单片机相关的头文件,剩下的 GEMM、归一化、采样逻辑完全可以在 x86 上跑。

然后做两件事:第一,把 PC 版的输出和 PyTorch 的输出逐 token 对比,一致就说明算法正确;第二,把板端的输出和 PC 版的输出对比,一致就说明硬件环境没问题。这样把一个"两端黑箱"的问题拆成了两个"单端白箱"的问题,定位效率提升好几倍。

我还额外加了一个"逐层 dump"的调试开关,把每层的输出前 16 个值打印成十六进制,两边一比就能精确定位到是哪一层开始发散的。第一次用这个方法,我五分钟就找到了问题:某个矩阵的转置搞错了,PyTorch 里是[in, out],我导出的时候按[out, in]排的序。

5.2 常见问题速查表

下面这张表是我这三天里实际遇到并解决的问题,以及一些顺着经验推断出来的高频故障:

现象可能原因排查方法
链接报 RAM 溢出权重数组漏了const,被复制到 RAM检查所有静态权重是不是const
无规律 HardFault栈太小,或被大数组挤爆栈从 1KB 提到 2KB,大缓冲区改静态
输出全是同一个字符KV Cache 没有正确写入dump 第 0 层 K、V 矩阵对比
输出前几个字正常,后面乱位置编码索引算错,或上下文越界检查 pos 是否超过ctx_len
输出全是<unk>词表映射反了,或 id 越界打印 token id 序列核对
每 token 耗时异常长Flash 等待周期没配,或printf阻塞检查 FLASH->ACR 和串口配置
PC 和板子结果不一致舍入规则不同、int8int32符号扩展问题统一舍入规则,检查char的符号性
频率越高越不稳定等待周期不足、电源去耦电容不够降频测试,检查最小系统板电源

这里我要特别说一下char的符号性问题。C 标准里char是不是有符号是实现定义的,ARM 编译器默认是无符号char,x86 上的 GCC 默认是有符号。如果你的代码里用了char来存 int8 数据,两端行为就会不一致,而且这种 bug 极难发现。结论:所有定点数据一律显式用int8_t,永远不要图省事写char

5.3 输出质量差:是模型问题还是量化问题,怎么区分

这个判断其实有很明确的方法。按下面的顺序排查:

第一步,把 PC 上的量化模拟跑一遍,如果生成质量就已经很差,说明是模型本身训练不到位或者语料不合适,跟部署无关,回去调训练。第二步,如果 PC 量化后正常但板子不正常,说明是 C 实现的问题,走 5.1 的对拍流程。第三步,如果没量化正常、量化后变差,那就是量化粒度和位宽的问题,从 per-tensor 改 per-channel,或者把关键层(比如输出投影)保持更高精度。

我一开始拿到乱码输出,第一反应是"模型太小了",差点回去重新训练。实际上是对拍之后发现 C 实现里有三个 bug。在没做对拍之前,不要动模型。

还有一点是关于生成结果的评判标准。一个 4 万参数的模型,指望它对答如流是不现实的。我给自己定的标准是:在它训练过的语料分布内,能生成语法正确、长度合理、偶尔有惊喜的片段。超出这个分布,它必然胡说。接受这个前提之后,评价标准就清晰了。


6. 收尾:几个值得提前做的决定

这个项目做到最后,我发现真正影响成败的不是那些技术细节,而是一开始做的几个决定。提前想清楚,能省下大量返工。

6.1 工程习惯上值得坚持的几件事

第一,先在 PC 上把量化版跑通,再上板子。板子的调试成本是 PC 的十倍以上,一个断点设置和重启烧录就够你等半分钟。第二,所有内存分配都用静态数组,不要用malloc。20KB 的堆在嵌入式上没有意义,动态分配带来的碎片和不确定性只会让你更难排查问题。第三,把所有跟硬件无关的代码单独放在一个文件夹里,方便你直接在 PC 上编译测试。我就是靠这个习惯才做到五分钟定位到转置错误的。

还有个小习惯:每次改动只改一个变量,改完立刻记录耗时和输出样本。我从-O0到 28ms 的这六步优化,每一步都留了记录,后来回头才发现,其中"归一化改定点"这一步贡献了 17ms,是所有改动里最大的。如果不记录,你根本判断不出该往哪个方向优化。

6.2 这块板子后面还能怎么玩

这个项目本身已经跑通了,但延伸方向挺多的。

第一个方向是给模型加保护。STM32F103C8T6 支持 Flash 读出保护,把权重这种花了时间训出来的资产保护起来是有意义的。可以配合芯片唯一 ID 做一个轻量的绑定校验,让权重只在特定芯片上生效。这块属于常规的固件保护范畴,配置起来不复杂,但要注意开启后调试口的行为变化,别把自己锁在外面。

第二个方向是引入 FreeRTOS 做任务划分。现在的推理是阻塞式的,串口发送会打断节奏。把推理、串口收发、以及可能的传感器采集拆成三个任务,用队列传递结果,整个系统的响应性会好很多。M3 跑 FreeRTOS 完全没压力,20KB RAM 里划 3KB 给任务栈足够用。

第三个方向是升级硬件。如果你已经把这个 4 万参数的模型玩明白了,下一步可以换个内核带 DSP 指令、Flash 和 RAM 更宽裕的芯片,把模型放大到 30 万参数这个量级。那时候你会明显感觉到生成质量的跃升,而且很多现在需要手工优化的地方,硬件直接帮你做了。

最后一个方向是换任务类型。语言模型只是序列建模的一种,同样的引擎稍微改改就能做别的:比如基于历史数据预测下一个传感器读数、把简短的命令序列补全成完整指令、或者做异常模式的检测。架构完全不用动,只要换语料重新训练,导出权重的流程一模一样。这也是我觉得这个项目最有意思的地方——你搭的不是一个玩具,是一套可以在 20KB 内存里跑序列模型的通用骨架。

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

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

立即咨询