把大模型压到 1.58 bit,是端侧推理里绕不开的方向。但大多数讨论停在 GPU 和手机 SoC 上,很少有人真的把它塞进一块几块钱的 MCU。这篇教程带你把这件"看起来不可能"的事拆成可执行的步骤:训练一个小尺寸的三元权重模型,把它打包进 ESP32-S3 的 flash,写一个整数累加推理内核,然后拼出 3~4 块板子的集群,实测吞吐、延迟和功耗。
这篇能做出什么
跟着走完,你会拿到三样东西:
1. 一份可烧录的 ESP32-S3 推理固件,权重以 -1/0/+1 三值存放,推理时只做 int8 加减法,不做浮点乘法;
2. 一个能跑起来的多节点集群,主机分发任务、从机计算并回传,拓扑可以选数据并行或层流水线;
3. 一组你自己测出来的数据:单样本延迟、集群吞吐随节点数的变化、空闲/推理/传输三种状态的功耗。
先说清楚边界。 BitNet 1.58-bit 的核心是权重被约束到 {-1, 0, +1}、激活量化到低比特。公开的 BitNet 模型参数量在十亿级,权重文件体积远超 ESP32-S3 的片内 SRAM 和常见外挂 Flash 容量(具体以芯片数据手册和模型卡为准)。所以这篇跑的是"同构但尺寸可控"的小模型:量化逻辑、打包格式、整数推理内核、集群通信方式都一致,只是规模缩小到 MCU 放得下。手法可以迁移,但这里的数字不要当成大模型的性能。
另外提醒一句:MCU 集群的总算力远低于任何一颗手机 SoC。做这件事的价值在于理解量化、内存布局和通信开销,而不是追求跑分。
前置条件清单
硬件
- ESP32-S3 开发板 3~4 块。带 PSRAM 的型号更从容,PSRAM 容量以官方数据手册为准。
- USB 数据线若干,一个能同时插上的 USB Hub。
- 独立 5V 供电。WiFi 发射的瞬时电流会让 USB 口电压跌落,直接导致板子重启。
- 可选:USB 电流计或 INA219 模块,用来测功耗。
- 可选:逻辑分析仪,排查时序问题很有用。
软件
- ESP-IDF,安装方式以官方文档当前版本为准。
- Python 3 与 PyTorch,用于训练和导出权重。
- 一块能平铺多块开发板的桌面,方便观察和接线。
前置知识
- 能读懂 C 的指针和数组运算。
- 会执行
idf.py build/idf.py flash/idf.py monitor。
第一步:先算内存账,再定模型规模
不要先写代码,先做算术。ESP32-S3 片内 SRAM 在几百 KB 量级(以官方数据手册为准),外挂 Flash 常见 8 MB 或 16 MB。三值权重每个占 2 bit,也就是说:
```
参数 26 万 × 2 bit ≈ 64 KB
参数 100 万 × 2 bit ≈ 250 KB
参数 400 万 × 2 bit ≈ 1 MB
```
看起来很美,但激活值、中间缓冲、栈、WiFi 协议栈都要占 SRAM。所以第一版建议定在这个量级:4 层、隐藏维度 128、序列长度 32 的小型 Transformer 或 MLP 混合结构,参数量几十万,权重放 flash,激活放 SRAM。
任务也选小的。适合入门的三个:
- 字符级分类(判断一个短字符串属于哪一类)
- 8 位以内整数加法的结果预测
- 关键词唤醒(一段定长音频分帧后分类)
选一个有明确标签、能自己造数据、单样本推理在毫秒级的任务就够了。
第二步:训练一个三值权重的小模型
核心是 BitLinear:前向时把权重按 absmean 缩放后四舍五入到 -1/0/+1,反向用直通估计器(STE)传梯度。
```python
import torch
import torch.nn as nn
import torch.nn.functional as F
class BitLinear(nn.Linear):
def forward(self, x):
w = self.weight
权重三值化:-1 / 0 / +1
scale = w.abs().mean().clamp(min=1e-5)
w_q = torch.clamp(torch.round(w / scale), -1, 1)
w_q = w + (w_q - w).detach() # STE,梯度照常回传
激活量化到 int8
a_scale = 127.0 / x.abs().max().clamp(min=1e-5)
x_q = torch.clamp(torch.round(x * a_scale), -128, 127) / a_scale
return F.linear(x_q, w_q, self.bias) * scale
```
用它堆一个小模型,训练到验证集不再提升即可。训练完记得统计一下权重里 0 的比例——零越多,后面推理时能跳过的乘法越多,实测延迟会明显下降。
```python
def zero_ratio(model):
total = zeros = 0
for m in model.modules():
if isinstance(m, BitLinear):
w = m.weight.detach()
scale = w.abs().mean().clamp(min=1e-5)
wq = torch.clamp(torch.round(w / scale), -1, 1)
total += wq.numel()
zeros += (wq == 0).sum().item()
return zeros / max(total, 1)
```
第三步:把权重导出成 C 头文件
三值只占 2 bit,四个权重塞进一个字节。约定映射:-1 → 0b00,0 → 0b01,+1 → 0b10。
```python
import numpy as np
def pack_ternary(w: np.ndarray) -> np.ndarray:
idx = (w.astype(np.int8) + 1).astype(np.uint8) # -1->0, 0->1, 1->2
flat = idx.reshape(-1)
pad = (-len(flat)) % 4
if pad:
flat = np.concatenate([flat, np.ones(pad, dtype=np.uint8)])
flat = flat.reshape(-1, 4)
return (flat[:, 0] | (flat[:, 1] << 2) |
(flat[:, 2] << 4) | (flat[:, 3] << 6)).astype(np.uint8)
def emit(f, name, arr, ctype="uint8_t"):
f.write(f"static const {ctype} {name}[] = {{\n")
for i in range(0, len(arr), 16):
chunk = ", ".join(f"0x{v:02x}" for v in arr[i:i + 16])
f.write(f" {chunk},\n")
f.write("};\n\n")
with open("model_weights.h", "w") as f:
for name, w in export_dict.items():
emit(f, name, pack_ternary(w))
```
关键点:所有数组加 static const,链接器会把它放进 flash 的 rodata 段,不占 SRAM。每层的缩放因子 scale 单独导出成 float 数组,别省这一步。
第四步:写整数推理内核
推理时每个输出元素就是一次点积,但乘数只有三种取值,可以退化成加减。
```c
#include <stdint.h>
// act: int8 激活,长度 n(需为 4 的倍数)
// packed: 每字节 4 个三值权重
static inline int32_t dot_ternary(const int8_t *act, const uint8_t *packed, int n) {
int32_t acc = 0;
for (int i = 0; i < n; i += 4) {
uint8_t b = packed[i >> 2];
for (int k = 0; k < 4; k++) {
int v = (b >> (2 * k)) & 0x3;
if (v == 1) continue; // 权重为 0,整项跳过
int8_t a = act[i + k];
acc += (v == 2) ? (int32_t)a : -(int32_t)a; // 2->+1, 0->-1
}
}
return acc;
}
```
再写量化激活、ReLU 之类的辅助函数,把每层的 dot_ternary 串起来,乘上该层的 scale。
```c
int64_t t0 = esp_timer_get_time();
int pred = run_inference(&ctx, input_buf);
int64_t t1 = esp_timer_get_time();
printf("pred=%d latency_us=%lld\n", pred, (long long)(t1 - t0));
```
在 menuconfig 里把 CPU 主频设到较高档位,编译时打开优化。跑通后先在串口终端确认输出和 Python 端的推理结果对得上,误差在量化允许范围内。
第五步:组装多节点集群
有两种拓扑,先想清楚要测什么。
数据并行(测吞吐)
每块板烧同一份固件,都装完整模型。主机把一批样本拆开,通过 ESP-NOW 分发给从机,从机算完把结果回传。吞吐理论上随节点数线性增长,直到通信成为瓶颈。
层流水线(测能放多大的模型)
把 N 层切成 N 段,每块板跑一段,激活值沿链传递。这样能放下比单板更大的模型,但延迟是各段之和加上通信往返——节点越多,单样本延迟越高。
ESP-NOW 的连接初始化大致如下:
```c
#include "esp_now.h"
#include "esp_wifi.h"
static const uint8_t PEER_MAC[6] = {0xAA, 0xBB, 0xCC, 0xDD, 0xEE, 0xFF};
void cluster_init(void) {
ESP_ERROR_CHECK(esp_netif_init());
ESP_ERROR_CHECK(esp_event_loop_create_default());
wifi_init_config_t cfg = WIFI_INIT_CONFIG_DEFAULT();
ESP_ERROR_CHECK(esp_wifi_init(&cfg));
ESP_ERROR_CHECK(esp_wifi_set_storage(WIFI_STORAGE_RAM));
ESP_ERROR_CHECK(esp_wifi_set_mode(WIFI_MODE_STA));
ESP_ERROR_CHECK(esp_wifi_start());
// 所有节点必须处于同一信道
ESP_ERROR_CHECK(esp_wifi_set_channel(1, WIFI_SECOND_CHAN_NONE));
ESP_ERROR_CHECK(esp_now_init());
esp_now_peer_info_t peer = {0};
for (int i = 0; i < 6; i++) peer.peer_addr[i] = PEER_MAC[i];
peer.channel = 1;
peer.ifidx = WIFI_IF_STA;
ESP_ERROR_CHECK(esp_now_add_peer(&peer));
}
```
两个必须记住的约束:ESP-NOW 单包负载有上限(以官方文档为准,通常两百多字节),层激活值大了要手动分片重组;收发回调的函数签名在不同 IDF 版本里改过,照着所用版本的 API 参考写。
节点 MAC 地址可以在启动时打印出来,手工抄进主机固件,或者加一段发现流程让从机主动上报。
第六步:测吞吐、延迟与功耗
延迟用 esp_timer_get_time() 打点,重复几百次,取中位数和 P95,不要只报平均值——平均值会被偶发的 WiFi 重传拉偏。
吞吐用 样本数 / 总耗时。测的时候先跑一批 warm-up,把 flash cache 预热和 WiFi 关联的抖动排除掉。
功耗三种状态分开测:
- 空闲:只跑主循环,不推理不通信
- 推理:连续推理,关掉 WiFi
- 传输:只做 ESP-NOW 收发
公式很简单:P = U × I,能量是 P × t。用 USB 电流计读取,或者用 INA219 通过 I2C 采回来记录。
建议做一张这样的表:
| 节点数 | 单样本延迟(ms) 中位数/P95 | 吞吐(样本/秒) | 推理功耗(mW) | 传输功耗(mW) |
|---|---|---|---|---|
| 1 | ||||
| 2 | ||||
| 3 | ||||
| 4 |
数据并行下你会看到吞吐增长逐渐放缓,那条曲线的拐点就是通信瓶颈开始的位置。层流水线下延迟随节点数上升,正好和吞吐趋势相反——这就是为什么"集群"不等于"更快",取决于你把任务切在哪一层。
常见坑与排错
一上电就重启。 大概率是供电。WiFi 发射的瞬时电流能到几百毫安,劣质 USB 线或 Hub 撑不住。换独立 5V 电源,或在电源脚旁边并一个大电容。
编译提示 rodata 溢出。 权重数组忘了加 static const,被当成可写数据放进了 SRAM。检查每个权重数组的声明。
推理结果和 Python 端差很远。 按顺序查三件事:打包时的位序(本教程用的是低位在前)、每层 scale 是否对上、激活量化的截断范围是不是 -128~127。写一个最小的 4 元素点积单元测试,在 PC 上用 C 编译跑一遍对比,比在板子上盲猜快得多。
任务看门狗超时。 一次完整推理如果超过几秒,会在长序列或大层上触发。要么分片推理、每片之间喂狗,要么在 menuconfig 里调大看门狗超时。
ESP-NOW 发不出去。 先确认两块板的信道一致,再确认 add_peer 成功返回、MAC 地址没抄错。同一个信道里如果还有别的 WiFi 在跑,丢包率会明显上升。
延迟数据忽高忽低。 串口日志本身就是开销。测量时把日志等级调到 warn 以上,或者把结果先存数组,测完再一次性打印。
PSRAM 用了但没变快。 确认 PSRAM 初始化成功、并且权重确实读自 flash 缓存或 PSRAM。片内 SRAM 访问速度高于外挂 PSRAM,把热点激活值留在片内。
下一步建议
跑通之后,几个值得继续推的方向:
扩大规模。 把节点加到 6~8 块,观察吞吐曲线的拐点往哪走,能反推出 ESP-NOW 的有效带宽。
换拓扑。 试环形全归约(ring all-reduce),对比星型拓扑在同样节点数下的通信开销。
压更狠。 在训练时加稀疏约束,提高权重里 0 的比例。三值里已经有一路是 0,把零的比例推高,dot_ternary 里的 continue 命中率更高,延迟直接下降。
用上硬件加速。 ESP-DSP 或 CMSIS-NN 里有点积和 SIMD 指令,把内核换过去,对比手写循环的差距。
接真实任务。 找一个你手边有的小任务,比如几个关键词的语音唤醒、传感器数据的短序列分类,把整条链路走通一次——从数据采集、训练、导出到固件部署。这件事做完,你对"端侧模型"的理解会比读十篇论文都具体。
最后提醒:芯片主频、PSRAM 容量、ESP-NOW 负载上限这些参数,以及 BitNet 模型的具体规模,都以官方文档和模型卡当前版本为准。数字会变,方法不变。
