结构化输出生成:约束解码与语法引导生成的工程实战
一、问题:大模型为什么"不会"写JSON
2024 年底,我在一个 LLM-as-a-Service 项目中遇到了一个看似荒谬的问题:GPT-4 写了整整一年的"JSON 输出",但每天仍有约 3% 的概率输出格式错误的 JSON——多余的逗号、丢失的引号、转义字符缺失。这不是模型能力问题,而是自回归生成机制与结构化输出之间的根本矛盾。
问题的本质在于:语言模型每次只生成一个 token,它"不知道"前面生成的 token 是否会导致最终结果不合法。就像一个盲人画圆,每一笔都正确,但最终无法保证闭合。更糟的是,即使我们在 prompt 中说"请输出合法的 JSON",模型也无法在 token 级别保证这一点——因为模型预测的是下一个 token 的概率分布,而非语法约束。
过去两年,工业界发展出了一套优雅的解决方案:约束解码(Constrained Decoding),也称为语法引导生成(Grammar-Guided Generation)。其核心思想简单而强大——在每一步 token 采样之前,屏蔽掉所有会导致最终输出违反语法规则的候选 token。这不是 prompt 技巧,而是在 logits 层面进行数学干预。
二、约束解码的数学原理
2.1 从概率分布到约束分布
标准自回归解码中,第 t 步的 token 选择为:
$$x_t \sim P(x | x_{ 其中 $P$ 是模型输出的 logits 经过 softmax 后的概率分布。约束解码对此进行修正: $$x_t \sim P(x | x_{ 其中 $\mathbb{C}(x, x_{ 这意味着: 约束解码有两种主流实现方式:logit masking 和 token rejection。 Logit Masking(掩码法) 是直接做法。在模型输出 logits 后、softmax 之前,将非法 token 对应的 logit 设置为 $-\infty$: Token Rejection(拒绝采样法) 则先按正常流程采样到一个 token,如果该 token 非法,则丢弃并重新采样,重复直到获得合法 token。这种方法更简单但效率低下——尤其在约束严格时(如需要输出特定格式的日期),可能需要数十次拒绝采样才能命中一个合法 token。工业界普遍采用 logit masking。 最简单的约束来源是正则表达式。例如,要求输出匹配 关键复杂度在于:一个 token(如 真正的挑战在于 JSON、YAML、XML 等结构,它们是上下文无关文法(CFG),无法用正则完整描述(正则无法处理任意嵌套)。JSON 的简化 EBNF 如下: 要为一个 CFG 构建约束解码器,需要下推自动机(PDA)。PDA 相比 DFA 多了一个栈(stack),栈用来跟踪嵌套结构——每进入一个 这是 Outlines、XGrammar 等框架的核心算法。 Outlines 的创立思想颠覆了传统观点——他们认为"约束解码不应该发生在 decode 循环的每个 step",而是可以交错(interleave)执行:让模型在不受约束时自由生成文本,只在需要结构化输出时切换到约束模式。 Outlines 的实现核心是将 CFG 编译为映射器(Mapper)——token 序列到 FSM 状态的映射。在每一步,它只遍历"当前 FSM 状态下仍可到达的 token",而不是遍历整个词汇表,这使得对大词表模型(128K tokens)也能高效运行。 XGrammar 来自陈天奇团队,代表了约束解码的工业级实现。其核心创新是编译时优化——将 CFG/正则预编译为高效的字节码解释器,避免运行时的逐 token 正则匹配开销。 XGrammar 的另一个杀手特性是适配器模式(Adapter)——它可以与不同推理引擎(vLLM、SGLang、llama.cpp)的 logits 处理管线无缝对接。其内部的"cache"机制会缓存常见状态的掩码,避免重复计算,可将约束开销控制在总生成时间的 5-10%。 Microsoft 的 Guidance 与上面两个截然不同。它不只是"约束输出",而是将"生成+约束"变成交互式的编程体验: Guidance 的独特之处在于:它把约束解码融入了"提示词模板语言"中,开发者像写 f-string 一样写结构化输出。 约束解码面临的核心性能挑战是每一步都需要生成一个 $O(V)$ 的掩码($V$ 是词表大小)。当 $V=128K$ 且模型在 4x GPU 上批量推理时,这一步的开销不可忽视。 工业实现中不会每次遍历全部 128K token,而是维护一个"当前合法 token 列表": 不同框架对这个循环的优化策略不同: 约束解码不能独立于采样策略存在。常见的交互问题包括: Top-p 采样(核采样):先屏蔽非法 token,再在合法 token 中做核采样。这保证了结果既符合约束,又保持多样性。 Beam Search:约束解码在 beam search 下更复杂——每条 beam 需要维护独立的状态。这是 vLLM 在支持结构化输出时 beam search 性能较差的原因。 即使有约束解码,生产环境中仍需处理异常。常见措施: 约束解码与流式生成(streaming)需要特殊处理。当用户希望边生成边看到结果时,约束解码必须支持部分掩码:在生成 Outlines 的解决方式是延迟 token 发射:当模型输出一个 token 时,先将其放入缓冲区,只有当该 token 确定是"不可撤销的路径上的必经 token"时才将其流式推送。 Agent 场景中,结构化输出通常用于工具调用(Function Calling)。一个请求可能触发多个并发的工具调用,每个工具的 schema 不同: vLLM 目前的实现会为每个请求维护一棵"约束森林",在多个并行 tool call 间共享同一 logits 处理器。 我们在生产环境(A100 80GB,LLaMA-3-8B,batch_size=8)上对约束解码的开销进行了基准测试: 可见,约束开销与 grammar 的复杂度正相关,但在最坏场景下也仅增加约 60% 的延迟——完全可接受。相比之下,使用 prompt + 重试的方式,一次失败后的 retry 开销通常是 100% 以上(需要重新生成平均 half 长度的输出)。 我们团队的实际数据:在一个 RAG + 工具调用的生产系统中,启用约束解码后: 约束解码领域正快速演进,几个值得关注的方向: 1. 推测解码(Speculative Decoding)的整合:draft model 本身也需要满足约束,如何让 draft 和 verify 阶段共享约束状态是当前研究热点。DeepMind 2025 年发布的"Taramana"框架已将约束融入推测解码。 2. 硬件加速 logits masking:Apple 和 Qualcomm 在 NPU 中开始支持"per-token masking"原语,未来约束解码可能成为硬件级操作。 3. 动态 grammar 切换:在长对话中,随着对话状态变化,可用工具集合也在变化,支持运行时无缝切换 grammar 的能力将成为 Agent 框架的基础设施。 工程建议总结: 约束解码代表了 LLM 应用工程的一个重要范式转变——我们不再乞求模型"请输出正确格式",而是直接保证输出的正确性。这不仅是体验问题,更是构建可信赖 AI Agent 系统的基础设施级别能力。
2.2 两种实现路径
import torch
def apply_logit_mask(logits: torch.Tensor, valid_mask: torch.Tensor) -> torch.Tensor:
"""
将无效 token 的 logit 设置为 -inf
Args:
logits: 形状为 [vocab_size] 的原始 logits
valid_mask: 形状为 [vocab_size] 的布尔掩码,True 表示该 token 合法
"""
# 将 False 位置设为 -inf
logits = logits.masked_fill(~valid_mask, float('-inf'))
return logits
# 示例:假设当前 only "name" 是合法下一个 token
valid_tokens = torch.zeros(128256, dtype=torch.bool) # LLaMA vocab size
valid_tokens[token_id_for_name] = True # 只允许 "name" token
# 应用掩码
masked_logits = apply_logit_mask(raw_logits, valid_tokens)
# 现在 softmax(masked_logits) 中非法 token 的概率为 0
三、从正则表达式到上下文无关文法
3.1 正则约束与有限状态机
\d{4}-\d{2}-\d{2} 的日期。合法的做法是将正则表达式编译为确定有限自动机(DFA),DFA 的每一步状态转移对应一个 token 分类。
# 伪代码:将正则编译为 DFA,并为每个状态维护合法 token集合
regex = r'\d{4}-\d{2}-\d{2}'
dfa = compile_regex_to_dfa(regex)
# 在解码循环中
current_state = dfa.initial_state
while not dfa.is_accepting(current_state):
# 获取当前状态允许的 token 集合
valid_tokens = dfa.get_valid_tokens(current_state)
# 创建 logit 掩码
mask = torch.zeros(vocab_size, dtype=torch.bool)
for token_id in valid_tokens:
mask[token_id] = True
# 应用掩码并采样
masked_logits = raw_logits.masked_fill(~mask, float('-inf'))
token_id = sample(masked_logits)
# 状态转移
current_state = dfa.transition(current_state, token_id)
"2024")可能对应多个字符,而 DFA 的状态转移是字符级的。因此需要"字符- token"映射层。这是所有基于正则的约束解码器的共同工程挑战。3.2 上下文无关文法与 JSON
object = '{' (pair (',' pair)*)? '}'
pair = string ':' value
value = string | number | object | array | 'true' | 'false' | 'null'
array = '[' (value (',' value)*)? ']'
string = '"' char* '"'
number = int '.'? frac? exp?
{' 压栈对象上下文,每遇到出栈弹栈。在每个 PDA 状态 + 栈顶 + 输入 token 的组合下,确定合法的下一个 token 集合。四、框架深度解析:Outlines vs XGrammar vs Guidance
4.1 Outlines:交错式生成范式
[自由文本] → [约束片段] → [自由文本] → [约束片段] → ...
import outlines
# 初始化模型
model = outlines.models.transformers("meta-llama/Meta-Llama-3-8B-Instruct")
# 方式1:正则表达式约束
generator = outlines.generate.regex(
model,
regex_str=r"[A-Z][a-z]+ is \d+ years old"
)
result = generator("Generate a sentence about Alice: ")
# 输出: "Alice is 25 years old" — 保证格式
# 方式2:JSON schema 约束
from pydantic import BaseModel
from typing import List
class Character(BaseModel):
name: str
age: int
skills: List[str]
generator = outlines.generate.json(model, Character)
result = generator("Create a character named Bob: ")
# 输出: {"name": "Bob", "age": 34, "skills": ["Python", "C++"]}
# 格式 100% 合法,失败率 0%
4.2 XGrammar:编译优化与适配器模式
import xgrammar as xgr
# 预编译 grammar(一次性开销,后续可复用)
grammar_compiler = xgr.GrammarCompiler(
xgr.TokenizerInfo.from_huggingface("meta-llama/Meta-Llama-3-8B-Instruct")
)
# 编译 JSON schema
compiled_grammar = grammar_compiler.compile_json_schema({
"type": "object",
"properties": {
"query": {"type": "string"},
"count": {"type": "integer", "minimum": 1},
},
"required": ["query", "count"]
})
# 创建 logits 处理器
logits_processor = xgr.LogitsProcessor(compiled_grammar)
# 在 vLLM 中使用
from vllm import LLM, SamplingParams
llm = LLM("meta-llama/Meta-Llama-3-8B-Instruct")
sampling_params = SamplingParams(
temperature=0.7,
max_tokens=256,
logits_processors=[logits_processor] # 注入约束
)
outputs = llm.generate(["Extract query and count from: show me 5 cats"], sampling_params)
# 输出: {"query": "cats", "count": 5}
4.3 Guidance:交互式 REPL 式约束
from guidance import models, gen, select
# 加载模型
llm = models.LlamaCpp("/models/llama-3-8b.gguf")
# 交互式构建输出
result = llm + "Q: 法国首都是什么?A: " + gen(
name="answer",
regex=r"[A-Z][a-z]+ " # 只允许 "Paris" 这种格式
) + "\nQ: 人口多少?A: " + gen(name="population", regex=r"\d+万?")
# result["answer"] -> "Paris"
# result["population"] -> "215万"
五、Logit Masking 的性能工程
5.1 词表过滤优化
class OptimizedMaskGenerator:
def __init__(self, tokenizer, root_state):
self.tokenizer = tokenizer
self.state = root_state
# 预计算每个 token 的字节表示(token 实际文本)
self.token_bytes = [
tokenizer.decode([i]) for i in range(tokenizer.vocab_size)
]
def get_valid_mask(self, state) -> torch.Tensor:
"""返回当前状态下的合法 token 掩码"""
valid_tokens = []
# 不是遍历所有 128K token,而是只检查"可能符合"的 token
# 利用 FSM 的转移表,每个状态直接关联到下一状态的有效 token
for token_id, token_text in enumerate(self.token_bytes):
if self.can_transition(state, token_text):
valid_tokens.append(token_id)
mask = torch.zeros(self.tokenizer.vocab_size, dtype=torch.bool)
mask[valid_tokens] = True
return mask
def can_transition(self, state, token_text: str) -> bool:
"""判断 token_text 在当前 state 下是否合法"""
# 利用预计算的 transition table,O(1) 查找
return state in self.fsm_transitions.get(token_text, set())
5.2 与采样策略的交互
def constrained_top_p_sampling(logits, valid_mask, p=0.9, temperature=0.7):
# 1. 应用约束掩码
masked_logits = logits.masked_fill(~valid_mask, float('-inf'))
# 2. 温度缩放
scaled_logits = masked_logits / temperature
# 3. 在合法 token 上做 top-p
probs = torch.softmax(scaled_logits, dim=-1)
sorted_probs, sorted_indices = torch.sort(probs, descending=True)
cumulative_probs = torch.cumsum(sorted_probs, dim=-1)
# 移除累计概率超过 p 的 token
sorted_indices_to_remove = cumulative_probs > p
sorted_indices_to_remove[..., 1:] = sorted_indices_to_remove[..., :-1].clone()
sorted_indices_to_remove[..., 0] = 0
indices_to_remove = sorted_indices[sorted_indices_to_remove]
masked_logits[indices_to_remove] = float('-inf')
return masked_logits
六、生产部署实战
6.1 错误恢复机制
class SafeStructuredGenerator:
def __init__(self, model, grammar, max_retries=3):
self.model = model
self.grammar = grammar
self.max_retries = max_retries
def generate_with_fallback(self, prompt: str) -> dict:
"""
带错误恢复的结构化输出
三层约束: 词法约束 → 语法约束 → 语义校验
"""
for attempt in range(self.max_retries):
# 第一层: 约束解码保证格式
raw_output = self.model.generate(
prompt,
logits_processor=LogitsProcessor(self.grammar)
)
# 第二层: 语义校验
try:
parsed = json.loads(raw_output)
if schema_validate(parsed, self.schema):
return parsed
except json.JSONDecodeError:
pass
# 如果失败,在 prompt 中加入上次错误信息(反射)
prompt += f"\n[上次输出有误: {raw_output[:100]}... 请修正]"
# 所有重试失败,返回降级结果
return self.fallback(prompt)
6.2 流式生成中的约束
"{" 时,模型已经知道接下来必须是 "\"name\"" 或 "\"age\"",但它不能向用户暴露这个信息——否则就失去了流式体验。
# 伪代码:约束流式生成
class StreamingConstrainedGenerator:
def __init__(self, model, grammar):
self.buffer = []
self.model = model
self.grammar = grammar
def stream(self, prompt):
state = self.grammar.initial_state
while True:
# 采样下一个 token
logits = self.model.forward(prefix + buffer)
mask = self.grammar.get_mask(state)
token_id = sample(logits.masked_fill(~mask, -inf))
# 关键:判断这个 token 路径是否"确定"
remaining_path = self.grammar.get_forced_path(state, token_id)
if remaining_path.length > 1:
# 路径有多条可选分支 → 流式发射 token
yield token_id
self.buffer.append(token_id)
state = self.grammar.transition(state, token_id)
else:
# 路径确定 → 可以缓冲,更快输出
self.buffer.append(token_id)
if self.grammar.can_emit(state):
for buffered_token in self.buffer:
yield buffered_token
self.buffer = []
state = self.grammar.transition(state, token_id)
6.3 多工具调用的链式约束
# 需要生成: [{"name": "get_weather", "arguments": {"city": "北京"}},
# {"name": "send_email", "arguments": {"to": "[email protected]"}}]
# 约束解码需要维护多个并行的状态
class MultiToolConstraint:
def __init__(self, tool_schemas):
self.tools = tool_schemas
# 构建选择文法:先选 tool name,再进入对应 schema
self.grammar = self._build_choice_grammar()
def _build_choice_grammar(self):
# 顶层结构
grammar = """
start: "[" function_call ("," function_call)* "]"
function_call: "{" '"name":' STRING '"arguments":' value "}"
STRING: "\\"" ("get_weather" | "send_email" | "search") "\\""
value: object | array | STRING | NUMBER | "true" | "false" | "null"
object: "{" [pair ("," pair)*] "}"
pair: STRING ":" value
array: "[" [value ("," value)*] "]"
%import common.NUMBER
%import common.WS
%ignore WS
"""
return grammar
七、性能基准与量化权衡
场景
约束解码延迟 (ms/tok)
无约束延迟 (ms/tok)
开销
简单正则(日期)
12.3
11.8
4.2%
JSON Schema(3 字段)
13.1
11.8
11.0%
嵌套 JSON(5 层)
14.5
11.8
22.9%
复杂 Schema(union 类型)
16.2
11.8
37.3%
智能体 tool_calls(5 tools)
18.7
11.8
58.5%
八、未来方向与工程建议

发表评论 取消回复