约束解码深读:把 mask 做到零开销之前,我们跨过了四道坎

沿一条主线拆 constrained decoding / grammar-based sampling:所有算法都在实现同一个「可延伸性 oracle」,四代论文(Outlines / DOMINO / XGrammar / SGLang / GAD)各推进一关——token-字符错位、mask 成本、确定性片段的浪费、分布扭曲。附伪代码骨架、亲手实测 96.4% 的确定性 token 占比、以及一个手算完的分布扭曲例子。

第一次真正跑通 Outlines 的 JSON schema 约束时,我盯着日志里的 FSM 状态推进看了近半小时。让我停在那里的不是「它有效」,是一个反直觉的事实:生成一份 300 token 左右的 JSON,模型真正被允许”自由选择”的位置只有几十个——其余绝大多数步的合法 token 集里只有一个成员。那一刻我意识到,约束解码真正在做的不是”限制模型”,而是把生成任务从「300 步采样」重新记账成「几十步真采样 + 若干段常量填充」。这正是 SGLang jump-forward 论文的全部价值观。

更早还有一次翻车更能说明为什么这件事难。我用一个字符级 regex 直接去掩 token 的 logits,模型跑到一半开始吐乱码。凌晨调到眼花才反应过来:]} 在很多 tokenizer 里是一个 token,但在语法里是两个符号;反过来一个字符串转义可能被切进多个 token。字符级自动机和 BPE token 边界根本对不上——我建的自动机在合法路径上把合法 token 拒掉了。这个坑吃掉了 Willard & Louf 那篇 Outlines 论文出现前的绝大多数早期实现。

Constrained decoding / grammar-based sampling 的所有工作,都可以沿一条主线读:它们在实现同一个 oracle——「给定当前前缀,哪些 token 能被延伸成一个合法完整序列」。四代论文的差异,是把这个 oracle 的成本、表达力、和分布保真度各推进了一步。

站内前置:约束解码在整体保障链中的位置见《大模型格式化输出五层保障》第 4 层;采样机制见《采样不是玄学:从 logits 到 temperature、top-k、top-p》。本篇聚焦第 4 层内部机制,做逐关拆解

名词速查

术语一句话解释
Constrained decoding每步采样前用一个”合法性 mask”把非法 token 的 logit 置 -\infty,物理上禁止非法输出
Grammar-based samplingConstrained decoding 的一种,合法性由形式文法(regex / CFG)定义
FSM / DFA有限状态机;判定正则语言,状态数有限,字符转移查表 O(1)
CFG / PDA上下文无关文法 / 下推自动机;能表达嵌套(递归 JSON、代码块),需要一个栈
BPE token大模型词表基本单元;一个 token 通常对应多字节,且切分不对齐语法符号
Vocabulary index预算好「FSM 状态 → 合法 token 集」的离线索引;运行时查表——Outlines 的核心机制
Jump-forward decoding当 FSM 处在确定性片段(下一 token 唯一)时,直接拼进序列、跳过前向
Extendability oracle本文用的抽象:判断”当前前缀 + 候选 token → 能否被延伸成合法完整序列”的函数
Logit mask一个词表大小的 bool 向量;True 位加 0、False 位加 -\infty,再走 softmax

一句话根本约束

约束解码之所以长成这样,是因为一条物理事实——LLM 是自回归的,一次采样一个 token;而语法合法性是对整段序列的全局判据。 这条鸿沟决定了:合法性只能被翻译成「每一步允许哪些 token」的逐步 oracle,而不能被翻译成”生成完再校验”。

反事实崩点:如果推理是并行/全局的(比如某类离散扩散模型),我们完全可以生成完再统一验收;但自回归下,任何非零的单步非法概率都会在长序列上以指数放大。粗算:每步 99.9% 合法、生成 1000 token,整体合法率 0.999100036.8%0.999^{1000} \approx 36.8\%——已经不能用。所以校验后重试的策略在长输出上会退化到主要靠”接着刷”,这就是为什么第 4 层不是可选项。

共同祖先:这套「每步用一个正确性 oracle 剪采样空间」的思路,最早的祖先其实是编译器里的 LR(k) 状态机——它也是一边扫、一边算”当前状态下合法的下一个符号集合”。约束解码是把这套东西逆向用在生成端。

四道工程门槛:主线

四代论文在解同一个 oracle 问题,各推进一关:

flowchart LR
    O["朴素做法<br/>每步跑一次 parser<br/>正确但慢到不可用"]
    T1["第 1 关<br/>Token-字符错位<br/>Outlines / DOMINO"]
    T2["第 2 关<br/>Mask 成本<br/>Outlines FSM 索引<br/>+ XGrammar CFG 分裂"]
    T3["第 3 关<br/>确定性片段浪费<br/>SGLang jump-forward"]
    T4["第 4 关<br/>分布扭曲<br/>Grammar-Aligned Decoding"]
    O --> T1 --> T2 --> T3 --> T4

结论摆前面,正文再逐关拆:

关口本质问题代表方案关键论文坏掉会怎样
Token-字符错位语法定义在字符上,生成走 tokenVocabulary index / 子词对齐Willard & Louf 2023 / DOMINO 2024合法输出被误拒、生成断流吐乱码
Mask 成本每步都算合法 token 集会拖垮吞吐编译期建索引 / 分裂 CI/CD tokenOutlines / XGrammar 2024每步 mask 开销 O(V·parse),吞吐掉一个数量级
确定性片段大量位置只有一个合法 token 却还在走前向jump-forward 直接拼常量SGLang 2023白算 90%+ 的前向,等于花大钱抽固定答案
分布扭曲局部 mask 后的联合分布 ≠ 真实条件分布ASAp 反馈迭代GAD 2024数据合成/评测采样偏差被系统性放大

第 1 关:Token-字符错位——最阴险的坑

语法是写在字符(或字节)上的,比如 JSON 的 {}";LLM 生成的是BPE token,一个 token 常常横跨多个字符。这道错位造成两类奇葩:

  • 一 token 多符号]}":}\n 在很多 tokenizer 里是单个 token——它一次跨过两个乃至更多语法符号;
  • 一符号多 token:字符串里的 \" 转义、Unicode 的 \uXXXX 可能被切进 2–3 个 token 的中间。

裸的字符级 FSM 不能直接掩 token 的 logits,因为它只知道”下一 char 是什么合法”,不知道”这个 token 走完会把 FSM 带到哪个状态”。

Willard & Louf 那篇 Outlines 论文(arXiv:2307.09702)给的解法直白但漂亮——编译期把字符级 FSM 提升到 token 级

# 伪代码,对应 Outlines 的 vocabulary index 构造
# 简化自 outlines/fsm/regex.py 的 create_fsm_index_end_to_end
def build_vocab_index(char_fsm, vocab):
    index = {}  # state_id -> {token_id -> next_state_id}
    for state in char_fsm.states:
        index[state] = {}
        for token_id, token_bytes in vocab.items():
            s = state
            ok = True
            for byte in token_bytes:
                if byte not in char_fsm.transitions(s):
                    ok = False; break
                s = char_fsm.transition(s, byte)
            if ok:
                index[state][token_id] = s   # 这个 token 在当前状态合法,走完到 s
    return index

运行时每步只需查 index[current_state] 拿到合法 token 集与其目标状态,掩码开销与词表大小几乎无关(词表是缓存了的,不是每步扫的)。这也是 Outlines 论文标题里 Efficient Guided Generation 的分量所在。

但这里藏了个更深的坑——即便自动机在 token 级正确,greedy 掩码依然会破坏分词。DOMINO(arXiv:2403.06988,Beurer-Kellner et al. 2024)观察到:如果 mask 强行把模型推去用”多个短 token”拼一个原本会用”单个长 token”表达的字符串,模型出来的分词分布就和无约束下的自然分词分岔了——perplexity 上升、下游质量下降。DOMINO 的做法是同时跟踪多种可行分词方案,结合投机解码把对齐做到近乎零开销、并在某些场景反而快 2×。它把”token-字符错位”这道坎从”能不能对齐”推到”能不能不破坏原生分词地对齐”。

我上文说”自己手搓九成 bug 出在这里”不是玩笑——这道坎的挖法有两层,第一层不难第二层难。

第 2 关:Mask 成本——从 O(V·parse) 到几乎零

一旦第 1 关解决,第 2 关就是把每步的掩码成本压到极致。

FSM 上是 O(1) 查表。Outlines 的 vocab index 每步就是一次哈希查询。

CFG 上就没这么便宜。JSON 里的对象嵌套({"a": {"b": {}}})不是正则能表达的——一个 pushdown automaton (PDA) 走这套语法要维护一个栈,“当前状态”就不再是有限的了,缓存爆炸。

XGrammar(arXiv:2411.15100,Dong et al. 2024)的解法是把词表切成两半

# 伪代码,对应 XGrammar 的 adaptive token mask cache
# 编译期一次性把每个 token 分类:
context_independent = set()   # 走完这个 token 不涉及栈变化,合法性只看 PDA state
context_dependent   = set()   # 合法性依赖当前栈内容(如"闭合括号"要看栈顶)

for token_id, token_bytes in vocab.items():
    if trace_pda_all_stacks(token_bytes).never_touches_stack():
        context_independent.add(token_id)
        # 离线为每个 PDA state 缓存 CI 部分的合法 token 集
    else:
        context_dependent.add(token_id)
        # 运行时用真实栈现算

# 运行时每步:
def legal_tokens(pda_state, stack):
    legal = ci_cache[pda_state]                # 查表,绝大多数 token 在此
    for token_id in context_dependent:         # 通常只有几百个,小 O(1)
        if simulate(token_id, pda_state, stack).ok():
            legal.add(token_id)
    return legal

论文报告:CFG 掩码的实际开销与 FSM 查表接近——这是 CFG 表达力第一次做到”和 regex 一样便宜”,vLLM 也因此把 XGrammar 作为默认的结构化后端之一。

判据:如果你的 schema 里有真的递归嵌套(对象套对象、数组套数组),选 XGrammar;如果你的 schema 是扁平枚举 + 固定字段(一层 dict + 若干 string/enum),Outlines 的 FSM 就够用且更简单。别用”我以后可能要嵌套”当理由——过度设计换不来真业务。

第 3 关:确定性片段——不采样也行

到了第 3 关,故事变得意外有趣。约束不再只是”额外开销”,反而变成”可以省事”的加速器。

观察:写死了固定字段的 schema,FSM 里有大量状态只有一个出边。等到那些位置,采不采样都是一样的结果,何必走一遍前向。

我自己跑了个玩具例子来量化这件事。取一个非常简化的枚举 schema:

{"action": ("approve"|"reject"|"escalate"), "score": ("low"|"mid"|"high")}

一次性验算脚本(python3 -c 直接跑):

import itertools
actions = ['approve', 'reject', 'escalate']
scores  = ['low', 'mid', 'high']
template = '{"action": "%s", "score": "%s"}'
paths = [template % (a, s) for a, s in itertools.product(actions, scores)]

def determinism(paths):
    total, forced = 0, 0
    max_len = max(len(p) for p in paths)
    for i in range(max_len):
        for pref in set(p[:i] for p in paths if len(p) > i):
            nexts = set(p[i] for p in paths if p.startswith(pref) and len(p) > i)
            total += 1
            if len(nexts) == 1: forced += 1
    return forced, total

f, t = determinism(paths); print(f, t, f'{100*f/t:.1f}%')

输出:

108 112 96.4%
{"action": "approve", "score": "low"}
============?===================?====

112 个位置里 108 个是被 FSM 完全钉死的,只有 2 处”真分叉”(那两个 ?:a/r/e 与 l/m/h 的首字符)。这就是 SGLang 那篇论文(arXiv:2312.07104)里 jump-forward 的直觉基础:别把 96% 的固定字符也走 forward pass,直接拼进序列、只在真分叉处让模型采样即可。论文里报告结构化解码在多个 benchmark 上的整体最高 6.4× 吞吐提升(含结构化解码的贡献部分)。

诚实提醒:这个 96.4% 是纯枚举 schema 的极端值。真实业务里如果 schema 里含自由文本字符串"description": <str>),字符串内部几乎每个字符都是”合法但非唯一”的,那段的确定性比例接近 0。所以 jump-forward 的红利与 schema 形态强相关——枚举/固定字段占比越高越赚。

顺带一提:jump-forward 的”一次前向多个 token”和投机解码(speculative decoding)在动作上是孪生兄弟——只不过投机解码的”草稿”来自小模型的猜测、需要主模型 verify;而 jump-forward 的”草稿”来自文法的确定性,天然可信,连 verify 都不需要。这条对比线在《Medusa 深读:解码慢不是算力问题,是搬运问题》里写过。

第 4 关:分布扭曲——最容易被忽略的一笔账

前三关都在解”效率/正确性”。第 4 关是保真度:即便 mask 完全正确,采样出的分布也不等于模型在”所有合法完整序列”上的真实条件分布。

Grammar-Aligned Decoding(arXiv:2405.21047,Park et al. 2024)把这件事讲透。用一个手算例子看清它。

设定:输出两个 token,每个是 ab;文法只禁止 aa。模型的真实分布是每步独立 P(a)=0.9,P(b)=0.1P(a)=0.9, P(b)=0.1

理论上”合法输出在真实分布下的条件概率”:

输出联合概率合法?
aa0.81✗ 禁
ab0.09
ba0.09
bb0.01

归一到合法集:P(ablegal)=0.09/0.1947%P(ab \mid \text{legal}) = 0.09/0.19 \approx 47\%ba:47%ba: 47\%bb:5%bb: 5\%

但逐步 mask 采出来的分布:第一步 ab 都还有救(都能延伸到某个合法后缀),不掩码——于是走 a 的概率是 0.9。走了 a 之后第二步 mask 掉 a,强制 b。结果:

输出mask 后概率
ab0.9×1.0=0.900.9 \times 1.0 = 0.90
ba0.1×0.9=0.090.1 \times 0.9 = 0.09
bb0.1×0.1=0.010.1 \times 0.1 = 0.01

ab 从 47% 被抬到 90%——贪心的局部掩码把「死路前最后一步」原本属于死路的概率错误地转给了同前缀的幸存者。模型看起来在自由采样,实际被文法拖着走了一条它其实不太想走的路。这个偏差在长序列、大分叉的场景会随路径长度累计。

GAD 的解法思路是引入一个修正因子,估计每个前缀下”能被延伸成合法序列的期望概率”,然后按这个因子重新加权 mask:

# 伪代码:ASAp 的核心,简化到只保留意图
# 对每个前缀 x_{<t},估计 A(x_{<t}) ≈ E[序列合法 | 前缀]
# 然后按 A 加权而不是简单 0/1 掩码

for step in range(max_len):
    logits = model(x)
    for token in vocab:
        if grammar.rejects(x + token):
            logits[token] = -inf
        else:
            # 关键改动:不是 +0,而是 + log A(x+token)
            logits[token] += log(estimate_A(x + token))
    x = sample(softmax(logits))

# A 的估计方法:迭代——先用普通 mask 采若干样,看合法率,作为 A 的初始值;
# 越采越准,直到收敛。

代价:需要多轮采样迭代逼近,吞吐显著下降。收益:分布真正对齐,数据合成/评测里不会因为文法而系统性偏斜。

判据:如果你的产出是喂给用户看的 JSON、function call——不用管,greedy mask 的偏差可以忽略;如果你在用 LLM 采样生成训练数据、或者用 LLM 做评测打分并要求分布无偏,就得盯 GAD 这条线。这条判据不写清楚容易吃大亏——评测数据被文法污染是最隐蔽的方法论 bug。

番外一关:能力代价(Let Me Speak Freely)

以上四关都在讨论”约束怎么落得干净”。真正独立的第五道账,是约束本身对模型能力的挤压

Tam et al. 2024 的 Let Me Speak Freely?arXiv:2408.02442)系统测量了这件事:在推理密集型任务上(数学、多步推理),越严格的格式约束,性能下降越明显。直觉解释:模型的”思考空间”就是它的 token 序列——一旦第一个 token 就必须落在 schema 里,它就没有”打草稿”的余地了。

工程上的标准缓解是两阶段

Stage 1: 无约束 CoT(让模型随便想)
Stage 2: 结构化提取(把 stage 1 的结论装进 schema)

或者更省事的做法:在 schema 里内置一个 reasoning 字段,把草稿纸画进表格里——OpenAI 的 Structured Outputs 官方指南就推荐这个做法。判据:任务越推理密集,越要给模型留 CoT 空间;纯提取/纯转换的任务,直接走 schema 无碍。

小结

一句话根本约束、崩点、可带走的动作压缩成 3 行:

  • 根本约束自回归采样是逐 token 局部决策,语法合法性是整段序列全局判据——差距只能靠”每一步一个 oracle”填。
  • 崩点:朴素做法每步跑 parser,1000 token 里合法率 0.999100037%0.999^{1000} \approx 37\%;同时 mask 成本 O(V·parse) 直接把吞吐打回一个数量级。
  • 可带走:选后端时问三个问题——(1) schema 有递归吗?没有用 Outlines/FSM,有用 XGrammar/CFG;(2) 静态字段占比大吗?大就开 jump-forward(SGLang);(3) 是拿来做数据合成/评测吗?是就查 GAD 那一条线的偏差。

收尾(可动手:5 分钟第一步)

装一次 outlines、跑一次它的 JSON 约束、把 FSM 状态数与 token-级索引大小打印出来——你会亲眼看到”96% 位置只有一个合法 token”这件事,从而用一次实验换到一次坐标:以后再看 SGLang / XGrammar / DOMINO 论文,你不再是从零理解,而是从”我在自己机器上见过的那个 FSM”出发。

pip install outlines transformers
python -c "
import outlines
from outlines.fsm.json_schema import build_regex_from_schema
schema = '{\"type\": \"object\", \"properties\": {\"action\": {\"type\": \"string\", \"enum\": [\"approve\",\"reject\",\"escalate\"]}}, \"required\":[\"action\"]}'
regex = build_regex_from_schema(schema)
print('Regex:', regex)
"

成本:装依赖 ~2 min,跑一次 ~30s。跑完你会拿到那个 regex,把它送进 outlines.processors 或 vLLM 的 structured outputs 就是完整的约束解码 pipeline——剩下的就是照着上面四关按需换后端。


诚实的提醒(这篇里比较软的地方)

  • 上面 96.4% 是我在纯枚举 schema 上一次性验算得到的极端值;真实业务的 schema 会因为字符串/自由文本字段把这个比例拉低到 60%–90%。原始验算命令与输出粘在了正文里,读者可以直接跑。
  • DOMINO 与 GAD 的伪代码是我按论文动机简化的意图版本,用来讲机制,不是原论文里的完整算法(DOMINO 的多分词跟踪、GAD 的 ASAp 收敛证明都要看原文才严谨)。
  • XGrammar 的分裂判据里 “never touches stack” 是我为读者简化的说法;原文用的是更严格的 PDA state 等价性论证。生产使用直接用它就是了,理解到这层已经足够读源码。

参考来源

arXiv 论文

  • FSM 约束解码:Willard & Louf, Efficient Guided Generation for Large Language Models (Outlines) — arXiv:2307.09702
  • 子词对齐 + 投机加速:Beurer-Kellner et al., Guiding LLMs The Right Way: Fast, Non-Invasive Constrained Generation (DOMINO) — arXiv:2403.06988
  • 高效 CFG 引擎:Dong et al., XGrammar: Flexible and Efficient Structured Generation Engine for LLMsarXiv:2411.15100
  • 压缩 FSM 与跳跃解码:Zheng et al., SGLang: Efficient Execution of Structured Language Model ProgramsarXiv:2312.07104
  • 分布扭曲与矫正:Park et al., Grammar-Aligned DecodingarXiv:2405.21047
  • 格式约束的能力代价:Tam et al., Let Me Speak Freely? A Study on the Impact of Format Restrictions on Performance of LLMsarXiv:2408.02442
  • 语言集成约束的先驱:Beurer-Kellner et al., Prompting Is Programming: A Query Language for LLMs (LMQL) — arXiv:2212.06094

工程实践

站内相关