Needle 2的QAT量化感知训练:fake_quant、STE与CQ噪声注入时机全解
【免费下载链接】needle14MB foundation model for tiny devices; phones, wearables, smart home, and robots.项目地址: https://gitcode.com/GitHub_Trending/needle20/needle
Needle 2 是一款面向手机、可穿戴设备等微型终端的 14MB 基础模型,而 QAT(量化感知训练)正是它能把 45M 参数压缩到 2-bit、全会话仅占 28MB 内存的核心技术。本文结合 needle/model/quantize.py 的真实实现,用零门槛的方式讲清 QAT 的三大主角——fake_quant仿真量化、STE 直通估计器、CQ 噪声注入——分别在什么时机生效、为什么这样设计,帮助新手快速看懂这套"边训练、边模拟量化"的完整流程。
为什么 Needle 2 必须靠 QAT:从 45M 参数到 14MB 二进制
先说结论:QAT 解决的是"低比特模型不聪明"这个老大难问题。
直接把 fp16 模型四舍五入到 2-bit(这叫"训练后量化",PTQ),精度往往崩塌。QAT 的思路则是在训练阶段就模拟量化误差,让模型在"带噪的世界"里学会生存,导出时再真正量化——精度损失被降到最低。
上图正是 QAT 的价值证明:Needle 2 以CQ2-bit(Cactus Quants 2 比特)量化后,在 Mobile-Actions 基准上打平 FunctionGemma 270M、LFM2.5 230M 这类 5~70 倍大的 fp16 模型——这不是靠硬件,而是靠 QAT 把量化损失"训没了"。
QAT 三件套:一张表看懂分工
在动手之前,先看全景。needle/model/quantize.py 中三个关键机制各司其职:
| 机制 | 角色 | 作用对象 | 生效时机 |
|---|---|---|---|
fake_quant | 前向"仿真量化",按 128 组 absmax 量化到 4-bit | 权重 kernel / embedding | 训练前向传播时,按QAT_EVERY频率触发 |
| STE(直通估计器) | 让梯度绕过不可导的 round 操作 | 所有仿真正/量化操作 | 反向传播时 |
| CQ 噪声注入 | 按 CQ 失真率标定强度,向权重加高斯噪声 | 可量化张量 | 训练时替代/配合硬量化 |
fake_quant:用"彩排"提前体验量化
核心函数在 needle/model/quantize.py#L10-L22,它做四件事:
- 分组:把权重按
group_size=128切成小组(维度不够就零填充再切); - 定标:每组取
absmax求缩放因子scale = absmax / qmax,qmax = 2^(bits-1) - 1(4-bit 时为 7); - 量化-反量化:
clip(round(w / scale)) * scale——这一来一回,连续权重被"掰"到最近的量化格点; - STE 缝合:最后返回
w + stop_gradient(q - w)。
第 4 步就是整个 QAT 最巧的一笔:前向传播看到量化后的值,反向传播看到的却是原权重——误差项(q - w)被stop_gradient冻结,梯度像没发生量化一样直接穿过。这就是所谓的STE(Straight-Through Estimator,直通估计器)。
为什么需要 STE?round()几乎处处不可导,如果梯度在这里断掉,量化前的权重永远收不到更新信号。STE 用一个"梯度恒等"的替身把路铺通:模型在带量化误差的景观里前向,却沿着平滑的原始路径回传梯度——两全其美。
同样的套路在部署侧也有镜像:cq_ste(needle/model/quantize.py#L354-L355)用w + stop_gradient(cq_quantize(w) - w)包住 CQ 码本量化,保证评估/混合精度路径同样可导。
CQ 噪声注入:注入什么、多少、何时
fake_quant模拟的是"格点量化",而 Needle 2 部署时用的是另一套CQ(Cactus Quants):Walsh-Hadamard 旋转变换 + Lloyd-Max 高斯码本(见cq_quantize,needle/model/quantize.py#L134-L147),甚至支持 1.58-bit 三值化。两种量化的失真特性不同,所以噪声注入必须"按部署方案定剂量":
- 剂量怎么算:
noise_scale(bits)先用cq_distortion在随机矩阵上实测 CQ 的相对失真率,再对比特数做对数插值,最终sigma = 组内 RMS × scale。也就是说:注入的噪声强度,精确等于"部署后真实会有的量化误差量级"; - 怎么加:
add_cq_noise(needle/model/quantize.py#L242-L251)对每个 128 组独立采样高斯噪声,noise_params则批量作用到所有可量化张量(kernel、embedding、mhc_phi)。
注入时机是 QAT 与普通训练的分水岭:
- 训练期:用噪声(或周期性
fake_quant)替代硬量化。权重每次更新都在"抖动",模型被迫学会对量化级误差鲁棒——这正是"量化感知"四个字的含义; - 部署期:噪声全部撤掉,
deploy_quantize直接执行真实的 CQ 量化,产出.cact归档。
权重量化本身的触发频率由configure_qat(every, ...)控制(needle/model/quantize.py#L56-L58):maybe_quant_weights通过jax.lax.cond按QAT_EVERY决定这一步是否真的过一遍 4-bit 组量化(组大小 128)。这种"间歇彩排"比每步都量化省算力,又足够让模型适应量化景观。
激活 8-bit 与 KV 缓存量化:推理路径上的注入点
除了权重,前向路径上的激活也被纳入 QAT。在 needle/model/architecture.py 中,_aq()(L22-L25)用jax.lax.cond(quant, fake_quant_act, ...)把 8-bit 激活量化(ACT_BITS = 8)精确地挂在四个位置:
- embedding 读出后;
- 每个注意力块内部(残差、输出投影);
- 最终 logits 之前;
- MTP(多 token 预测)分支的拼接处。
KV 缓存同样有一处"量化闸门":maybe_quant_kv(needle/model/quantize.py#L41-L45)在KV_BITS配置非零时,对 K/V 做 CQ 64 组 fake 量化——这是 256-token 滑窗推理时内存控制的关键。而configure_deploy会在比特配置变化时自动jax.clear_caches(),避免 JIT 缓存串味。
如上图所示,QAT 的量化点正好覆盖每个 Transformer 块的激活主干——权重 4-bit(组量化)、激活 8-bit、KV 可选 8-bit,三者共同构成部署时的完整量化配方,也即导出信息里常见的W4A8 / CQ W4A8标记。
从训练到部署:两套"仿真"如何闭环
整个生命周期的衔接在 needle/model/finetune.py 的build_main里完成:
- 微调:
needle finetune用 LoRA 微调冻结的基座(默认 rank 16 / alpha 32,AdamW + warmup-cosine 调度),细节见 doc/finetuning.md; - 构建:
needle build合并 LoRA 适配器,再按检查点声明的逐层比特映射(mixed bits,如敏感层 3-bit、其余 2-bit)走cq_quantize_params真实量化,输出单个.cact; - 部署:14MB 引擎直接加载
.cact,训练时的fake_quant/噪声仿真与部署时的 CQ 量化由同一套cq_quantize实现保证"仿真误差 = 真实误差"。
这也是 QAT 与 PTQ 的本质区别:PTQ 训完才量化,误差无处可逃;QAT 训时就量化(或等价噪声),误差被模型内化。
新手快速上手:三步跑通 QAT 微调流程
pip install cactus-needle needle finetune data.jsonl --epochs 3 needle build checkpoints/needle2.pkl --lora checkpoints/needle_lora.pkl --out my.cact- 数据格式:JSONL,一行一个
{query, tools, answers}示例(doc/finetuning.md 有完整说明); --bits 2|4或检查点内嵌比特映射控制量化宽度,--upload可发布归档;- 运行时用
needle.Needle(weights="my.cact", tools=[...])加载,引擎对权重无感知; - 推理行为契约与 API 细节见 doc/apis.md;权重加载机制的回归测试可参考 tests/test_weights.py。
总结:三个关键时机,一句话记住
- 前向时——
fake_quant按 128 组 4-bit absmax 把权重"掰"到量化格点,激活走 8-bit; - 反向时——STE 用
stop_gradient缝合,梯度无视量化墙直接穿过; - 训练全程——CQ 噪声按"部署失真率"标定注入,让模型在噪声中学出鲁棒性;部署时撤掉噪声、执行真量化,零落差上线。
这三件事缺一不可:没有 fake_quant,模型没体验过量化;没有 STE,梯度断流训不动;没有 CQ 噪声,仿真与部署之间就隔着一条精度鸿沟。读懂 needle/model/quantize.py 这不到 400 行代码,你就掌握了微型端侧模型 QAT 的完整方法论。
【免费下载链接】needle14MB foundation model for tiny devices; phones, wearables, smart home, and robots.项目地址: https://gitcode.com/GitHub_Trending/needle20/needle
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考