跳到主要内容

数据截至 (上游 commit f775db03aaa8)

04 · 结构化输出:在采样器上强制合法

本章讲什么: 「模型必须输出合法 JSON / 符合某条 regex」这个需求,怎么从「求模型自觉」变成「非法 token 根本采不到」。机制全在 python/sglang/srt/constrained/ + 采样路径上。

4.1 它要解决的小问题

工具调用、JSON 模式、结构化抽取这类场景,输出格式错一个字符,下游解析就崩。靠 prompt 工程「请严格输出 JSON」只是统计上更可能,不是保证。

真正的保证只有一个办法:在每一步采样时,把会导致非法输出的 token 直接从候选集里抹掉。

4.2 思路:grammar → FSM → 每步 bitmask

三步走:

  1. 编译。 把 json_schema / regex / EBNF 编译成一个有限状态机(FSM):每个状态知道「此刻哪些 token 合法」。
  2. 掩码。 每个 decode 步,FSM 给出当前状态的合法 token 集合,编码成 vocab 大小的 bitmask;把非法位置的 logits 置为 -inf,采样器自然采不到。
  3. 推进。 采出 token 后,FSM 吃下这个 token 进入下一状态,下一步的合法集合随之更新。

这套机制的两个工程难点,SGLang 都有专门设计:

  • 编译慢(regex/schema 编译是 CPU 重活,可能几十毫秒)——异步化,不挡调度循环。
  • TP 多卡下所有 rank 的调度决策必须一致——编译完成状态要跨 rank 同步。

4.3 一条带 grammar 的请求的一生

① 请求带 json_schema 到达


② process_req_with_grammar:查缓存;未命中 → 线程池异步编译,req.grammar = Future
│ 请求挂进 grammar_queue(不进 waiting_queue)

③ 调度循环每步轮询 get_ready_grammar_requests
Future 完成 + 跨 TP rank all_gather 对齐 → 移入 waiting_queue


④ 正常 prefill / decode,但每步采样前:
update_regex_vocab_mask 为每个带 grammar 的请求填一行 bitmask


⑤ apply_logits_bias:bitmask 打到 logits 上(triton kernel,非法 → -inf)


⑥ 采样出 token 后 advance_grammar_fsm / _accept_grammar_tokens 推进 FSM
grammar 自然终止(如 JSON 闭合)→ 请求结束

4.4 真实实现

异步编译与排队

GrammarManager(python/sglang/srt/constrained/grammar_manager.py:26)由 scheduler 在初始化时创建(scheduler.py:2007-2008)。每个带 grammar 的请求先到 process_req_with_grammar(grammar_manager.py:132,在 scheduler.py:2754 被调):

  • (key_type, key_string) 查 grammar 缓存,命中则 copy() 一份直接用。
  • 未命中则 ThreadPoolExecutor.submit 异步编译(get_cached_or_future_value,base_grammar_backend.py:284;线程池在 base_grammar_backend.py:206),req.grammar 暂时是一个 Future,请求挂进 grammar_queue

每轮组 prefill 批之前,_get_new_batch_prefill_raw 先查 grammar 队列(scheduler.py:3305-3309):get_ready_grammar_requests(grammar_manager.py:184)轮询各 Future,然后用 all_gather_object 在 PP0 的 DP/TP 组内取「所有 rank 都编好」的交集,再向后续 PP rank 传播——因为所有 rank 必须对「这条请求能不能进 waiting 队列」达成完全一致,否则组批发散、集群 hang。编译还有超时兜底:轮询超过 SGLANG_GRAMMAR_MAX_POLL_ITERATIONS 次直接判失败。

编译:xgrammar 后端

默认后端是 xgrammar(arg_groups/serving_hook.py:222-227,grammar_backendNone 时声明为 "xgrammar")。XGrammarBackend.dispatch_json(python/sglang/srt/constrained/xgrammar_backend.py:337)把 schema 交给 grammar_compiler.compile_json_schema,产出一个编译好的上下文;regex / EBNF / structural_tag 各有对应入口(dispatch_regexdispatch_ebnfdispatch_structural_tag)。编译失败不炸服务——返回 InvalidGrammarObject,请求带错误信息正常 abort。

掩码:每步的合法 token 集合

采样前,ModelRunner._preprocess_logits(python/sglang/srt/model_executor/model_runner.py:1805-1822)调 SamplingBatchInfo.update_regex_vocab_mask(python/sglang/srt/sampling/sampling_batch_info.py:240):

  1. 给整个 batch 分配一块 vocab 大小的 bitmask 张量。
  2. 只对「有 grammar、未结束、未终止」的请求逐行填 mask(fill_vocab_mask_batched)——其余行保持全合法,不花冤枉钱。
  3. 搬到 GPU,包成 GrammarMask

随后在 apply_logits_bias(sampling_batch_info.py:300)里 grammar_mask.apply(logits),xgrammar 后端落到 apply_token_bitmask_inplace_triton(xgrammar_backend.py:131)把非法位置写成 -inf。注意执行顺序:penalty 等 logit 变换分 pre/post 两段,bitmask 夹在中间(_apply_pre_grammar_logits_transforms → mask → _apply_post_grammar_logits_transforms),保证 grammar 约束不被后续变换破坏。

推进与终止

采出 token 后,BatchResultProcessor.advance_grammar_fsm(python/sglang/srt/managers/scheduler_components/batch_result_processor.py:799)让 FSM 吃下本步 committed 的 token;_accept_grammar_tokens(batch_result_processor.py:768)负责「吃到 grammar 终止符就截断」——spec decode 一次接受多个 draft token 时,越过终止符的尾巴直接丢弃,不进 KV、不发给客户端。

投机采样 + grammar 的 overlap 组合里,FSM 推进被挪进 verify() 内部做(grammar barrier,scheduler.py:1906_advance_pending_grammar),让 CPU 的 FSM 推进和 GPU 的 verify forward 重叠。

4.5 原理演示

# 示意,非源码:约束解码的一步
fsm = compile_json_schema(schema) # ① 编译(异步做,不挡调度)
state = fsm.initial_state()
while not state.is_final():
allowed = state.allowed_token_bitmask() # ② 当前合法集合,vocab 大小的 0/1
logits = model.forward(context) # GPU 前向
logits[~allowed] = float("-inf") # 非法 token 打成 -inf
token = sample(logits) # 正常采样(温度/top-p 照旧)
state = state.accept(token) # ③ 推进 FSM

重点看: 约束解码不改变采样算法本身,只是在 logits 上加了一道「合法性格栅」。温度、top-p 等参数照常工作,只是候选集被 FSM 预先筛过。

4.6 关键细节与坑

  • jump-forward(压缩 FSM)接口仍在,主路径已不走它。 2024-02 博客宣传的「compressed FSM 跳步解码」(遇到确定性片段一次跳多个 token)以 try_jump_forward 接口留在所有后端上(xgrammar_backend.py:164outlines_backend.py:80),outlines 的实现和说明在 python/sglang/srt/constrained/outlines_jump_forward.py(文件头注释仍引用那篇博客)。但在本 commit 的 scheduler 解码路径里搜不到任何调用方——当前主路径是每步 bitmask + overlap。这是从代码读出的现状,不是推测:想用跳步加速的人需要知道这一点。
  • 编译缓存按字符串精确匹配。 key 是 (key_type, key_string) 元组(base_grammar_backend.py:284),schema 里多一个空格都是新 key——客户端最好规范化 schema 文本。
  • 含 NUL 字节的 grammar 直接拒编。 _grammar_key_contains_nul(base_grammar_backend.py:162)——防止二进制串键污染缓存与日志。
  • grammar 位掩码张量用完即释放。 _preprocess_logits 里打完 mask 立刻把 sampling_info.grammar_mask = None(model_runner.py:1788-1792 的注释):overlap 模式下闭包和 batch_record_buf 会延长 tensor 寿命,不手动释放会稳定漏显存。
  • 「reasoner + JSON」有特殊处理。 思考模型的「先想再答」用 ReasonerGrammarObject(reasoner_grammar_backend.py)包住真正 grammar,思考段不约束、答段才约束。
  • 只保形式,不保语义。 bitmask 保证输出被 grammar 接受;字段内容是否合理仍是模型的事。

4.7 本章代码地图

主题文件路径符号名
grammar 管理器python/sglang/srt/constrained/grammar_manager.pyGrammarManager.process_req_with_grammarget_ready_grammar_requests
后端抽象与缓存python/sglang/srt/constrained/base_grammar_backend.pyBaseGrammarBackendget_cached_or_future_value
xgrammar 实现python/sglang/srt/constrained/xgrammar_backend.pyXGrammarBackend.dispatch_jsonXGrammarGrammar.accept_token
outlines 跳步python/sglang/srt/constrained/outlines_jump_forward.pyOutlinesJumpForwardMap(历史路径)
掩码生成python/sglang/srt/sampling/sampling_batch_info.pySamplingBatchInfo.update_regex_vocab_maskGrammarMask.apply
logits 处理顺序python/sglang/srt/model_executor/model_runner.pyModelRunner._preprocess_logits
FSM 推进python/sglang/srt/managers/scheduler_components/batch_result_processor.pyadvance_grammar_fsm_accept_grammar_tokens