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() 的策略:

  1. 先数 prompt 里有多少个语义位置(:802-812);
  2. 多于 1 个语义位置、且后端要求单 token 语义 prefill 时,块大小 = 1(:814-816)——目前只有 CUDA 命中(backend_requires_single_token_semantic_prefill()),规避量化嵌入/多语义 prefill 的不稳定;
  3. 否则块大小取 clamp(128/n_gpu_layers, 8, 64):GPU 层越多块越小,控制每图的物化规模(:819-827);
  4. 分块循环逐块 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

This site uses Just the Docs, a documentation theme for Jekyll.