Slow-AR Transformer 计算图与 KV Cache
本章逐段拆解 eval_cached()——Slow-AR 唯一的计算实现。prefill()、step()、prefill_fast() 三个公开方法最终都调用它,区别只在于一次喂多少个 token:
| 公开方法 | n_tokens | 用途 |
|---|---|---|
| prefill() | 多 token,可分块 | 提示段预填充 |
| step() | 1 | 生成循环中每帧前进 |
| prefill_fast() | 任意 | 不分块的直接评估(generate() 用它做 prompt prefill) |
1. 入口校验与 head_dim 推断
:863-875 先做三道校验:token 数为正、输入长度恰好等于 n_tokens × codebook_dim(11)、n_past_ + n_tokens 不超过 KV cache 容量。任何一项不过都直接返回 false——这保证后续张量偏移计算不会越界。
head_dim 的推断值得留意(:881-886):带 QK-norm 时从 q_norm 权重形状取(ne[0]),否则用 wo 输出列数除以头数。这比直接用 dim/n_head 更稳健,能适应非标准 head_dim 的变体。随后算出:
q_size = 32 × head_dim = 2560 kv_size = 8 × head_dim = 640
attn_scale = 1/√head_dim sem_scale = 1/√11
2. 主机侧输入准备:掩码是核心
:903-921 的循环把扁平输入拆成建图所需的若干向量,关键是区分语义位置与非语义位置:
const bool is_semantic = semantic ∈ [semantic_begin, semantic_end];
semantic_vals[t] = semantic; // 第 0 行原值
pos_vals[t] = n_past_ + t; // 全局位置(RoPE 用)
semantic_mask_vals[t] = is_semantic ? 1.0 : 0.0; // 码本贡献开关
cb_vals[cb][t] = code_id + cb*codebook_size; // 仅语义位置有效
语义位置(参考音频的已知码段、以及生成回灌的帧)才有真实码本 id;普通文本位置掩码为 0,码本行被完全屏蔽。scale_codebook_embeddings 开启时还会生成逐 token 的缩放向量(语义位置用 sem_scale,其余 1.0,:912-914)。
3. 上下文复用与图的创建
ctx_buf_.resize(10 * 1024 * 1024); // 首次分配后复用(:923-926)
ggml_context * ctx0 = ggml_init({ctx_size_, ctx_buf_.data(), true});
ggml_cgraph * gf = ggml_new_graph_custom(ctx0, 32768, false);
10 MB 原始内存缓冲作为图上下文的存储,每次调用从同一缓冲重新 ggml_init——前一次的图自然作废,无需手动释放,也避免反复 malloc。输入 id 张量(语义 id、位置、掩码、各码本 id)先建空壳,数据在图分配完后才 tensor_set 填入(:1071-1079),这是 ggml 推荐的”输入张量复用”模式。
4. 嵌入层:语义 + 10 码本求和
x = get_rows(embeddings, semantic_ids) (:941)
codebook_sum = Σ_cb get_rows(codebook_embeddings, cb_ids) (:946-952)
x = x + codebook_sum * repeat(semantic_mask) (:954-959)
要点:
- 10 个码本共用一张 40960 行(10×4096)的嵌入表,靠
cb * codebook_size偏移寻址; - 每次
get_rows的结果逐个ggml_add累加,再用semantic_mask广播相乘——非语义位置整个码本和归零; - 量化嵌入的查询结果统一
cast成 F32(:942、:950),后续层全部在 F32 上运算。
5. 单层计算图(:964-1050)
36 层循环,每层结构完全一致。以第 il 层为例:
5.1 注意力前处理
attn_in = RMSNorm(x, attention_norm) rms_norm_weighted(:46)
qkv = attn_in × wqkv^T 融合 QKV 一次矩阵乘
q2d/k2d/v2d = qkv 的三个视图(偏移 0 / q_size / q_size+kv_size)
q,k,v = reshape_3d(head_dim, n_head[_kv], n_tokens)
视图切分在 :971-977,不拷贝数据;q2d 等先 cont 是因为视图行跨度与 3D 重塑要求不兼容。若启用 QK-norm,对 q、k 再各做一次 RMSNorm(:979-982)。
5.2 RoPE
q = rope_ext(q, positions, ..., base=rope_freq_base, ...) (:984)
k = rope_ext(k, ...) (:987)
位置直接使用准备好的全局位置 n_past_+t,RoPE 参数(基频 1e6、上下文长度)全部来自 hparams,不做线性/NTK 插值。
5.3 写入 KV cache
这一层是理解 cache 布局的关键(:991-1005):
层偏移 = il × memory_k_->nb[3] (每层一段)
token 偏移 = n_past_ × nb[2] (本批起始槽位)
k_slot = memory_k_ 的 3d 视图(head_dim, n_head_kv, n_tokens)
cpy(k → k_slot) / cpy(v → v_slot) 显式加入计算图
两条 cpy 是图里仅有的副作用节点:本批新算的 k/v 被追加写入持久 cache 的正确层、正确槽位。
5.4 拼接历史、GQA 扩展
若 n_past_ > 0:
k_past/v_past = memory_* 中前 n_past_ 槽位的视图 → reshape_3d
k_mem = concat(k_past, k) v_mem = concat(v_past, v) (:1018-1019)
k_rep = repeat_interleave_heads(k_mem, n_head/n_head_kv=4) (:1025)
v_rep = repeat_interleave_heads(v_mem, 4)
repeat_interleave_heads() 把 8 个 KV 头各复制 4 份扩展成 32 头,完成 GQA。F16 cache 在拼接前按需 cast 回当前类型(:1016-1017)。
5.5 注意力本体
KQ = K^T × Q permute 后 mul_mat(:1028-1030)
KQs = KQ × attn_scale
KQm = diag_mask_inf(KQs, n_past_) 因果掩码:本批第 i 行只能看前 n_past_+i 列(:1032)
KQf = softmax(KQm)
KQV = V × KQf
diag_mask_inf 的第二个参数是”历史长度”——分块 prefill 时每块都要让掩码从 n_past_ 开始,这正是该函数支持偏移掩码的用途。输出经 permute 后 cpy 到一块新的 F32 张量 attn_cur(:1038-1039),再投影 wo。
5.6 残差与 SwiGLU FFN
h = x + attn_out
gate = RMSNorm(h) × w1^T
up = RMSNorm(h) × w3^T
ff_h = swiglu_split(gate, up) ggml 融合算子 = silu(gate)⊙up(:1046)
x = h + ff_h × w2^T
使用融合的 ggml_swiglu_split 而非手工 silu + mul,减少一次中间物化。
6. 输出:hidden 与 tied logits
36 层结束后(:1052-1059):
slow_out = RMSNorm(x, norm)
hidden_last = last_token_view(slow_out) 的副本 形状 (dim, 1) = 2560
logits = embeddings × hidden_last tied:复用输入嵌入矩阵
last_token_view()(:70)取本批最后一个位置——单步时就是唯一位置,prefill 时是序列尾;- logits 的
mul_mat直接拿weights_.embeddings当权重,输入输出嵌入严格共享,与导出脚本的 tied 设定一致; hidden_last还要在下一章被 Fast-AR 使用,所以必须物化(cpy 到独立缓冲)。
7. 调度、计算、取回
sched_reset → sched_alloc_graph(gf) 跨后端分配,失败即返回(:1062-1069)
tensor_set 填入所有输入 (:1071-1079)
sched_graph_compute (:1081)
tensor_get 取回 hidden(2560) + logits(155776)(:1088-1091)
n_past_ += n_tokens (:1095)
计算结束后 result 同时携带两样东西:供采样语义 token 的 logits,和交给 Fast-AR 的 hidden。这就是双 AR 的交接点。每次调用前后都 sched_reset,让调度器可以为不同 n_tokens 重新分配跨后端缓冲。
8. KV cache 的组织与生命周期
init_kv_cache() 在每次合成开始时调用,按合成实际长度(prompt 列数 + max_new_tokens)分配,而不是一上来就开 32768:
memory_k_/memory_v_ = 4d(F16, head_dim, 8, max_seq_len, 36)
- 形状维度依次为:head_dim、KV 头数、序列长度、层数;
- 只要
n_gpu_layers_ > 0,整块 cache 都分配在 GPU(:749-750)——这就是 README 中”KV cache 是权重之外最大开销”(约 2.5 GB)的来源,与卸载层数无关; - 分配后立刻 memset 清零(:756-757)。
cache 在合成结束时由 clear_kv_cache() 显式释放(buffer + context 一起),所以同一条 pipeline 连续合成多次不会累积显存;每次合成前重新 init,长度按需伸缩。
9. 自适应分块 prefill
为什么不一次性把整个 prompt 建成一张大图?因为参考音频段可能有几百上千帧,单张图的中间物化会撑爆显存/内存。prefill() 的策略:
- 先数 prompt 里有多少个语义位置(:802-812);
- 多于 1 个语义位置、且后端要求单 token 语义 prefill 时,块大小 = 1(:814-816)——目前只有 CUDA 命中(backend_requires_single_token_semantic_prefill()),规避量化嵌入/多语义 prefill 的不稳定;
- 否则块大小取
clamp(128/n_gpu_layers, 8, 64):GPU 层越多块越小,控制每图的物化规模(:819-827); - 分块循环逐块
eval_cached,KV cache 在块间通过n_past_自然衔接(:841-849)。
注意 generate() 里实际用的是 prefill_fast(不分块),分块策略主要服务 Pipeline/服务端的长参考音频场景。
10. 小结
Slow-AR 的实现是一张”重建节点、持久槽位”的图:每层的算子每次重建,但 KV cache 作为持久张量跨调用存在;掩码让文本位置与音频位置在同一套权重下统一处理;最后的 tied logits 和 hidden 同时服务于采样与 Fast-AR。第 06 章追踪交接之后的故事:Fast-AR 怎么用 hidden 解码 10 个码,以及 generate() 如何把两个 AR 串成完整主循环。
关键文件
| 位置 | 职责 |
|---|---|
| s2_model.cpp:860 | eval_cached 计算图主体 |
| s2_model.cpp:941 | 语义 + 码本嵌入 |
| s2_model.cpp:991 | KV cache 槽位写入 |
| s2_model.cpp:1025 | GQA 头扩展 |
| s2_model.cpp:1046 | SwiGLU FFN |
| s2_model.cpp:1058 | tied logits |
| s2_model.cpp:717 | KV cache 分配 |
| s2_model.cpp:793 | 自适应分块 prefill |