Skip to main content

06 - 接一个新模型

前面四篇都是在讲框架自己怎么运转。这一篇换个视角:你要接一个新 TTS 模型进来,具体要写什么。

假设昨天开源了一个新模型,结构是标准的三段 —— 文本进来先做预处理,中间一个自回归骨干吐码本,最后一个声码器还原波形。你想让它能通过 /v1/audio/speech 提供服务。

好消息是:你只需要写自己那个模型目录下的文件,框架层一行不用改,只剩两处必须手工连线。

坏消息是「接得进去」和「接得对」是两回事。真正让人翻车的三类问题都不在主流程上:

长什么样后果
请求边界HTTP 层替你把没填的采样参数填上默认值,而那套默认值是照着另一个模型调的你分不清这个值是用户要的还是 HTTP 层随手填的
张量在哪预处理图省事写个 .cpu(),自回归那边又得原样搬回显存每个请求白白多一次同步和两次拷贝
中止清理请求被取消有三条不同的时机,其中一条是「预处理算完了才发现请求早没了」少堵一条就是缓慢的内存泄漏

下面按实际动手的顺序走一遍。

一、要写的文件

模型代码全部关在 sglang_omni/models/<name>/ 里,框架层一行都不用碰:

models/<name>/
├── config.py # PipelineConfig 子类 + StageConfig 列表 + EntryClass
├── stages.py # 每个阶段的工厂函数
├── request_builders.py # 阶段间的载荷变换
├── payload_types.py # 该模型专有的流水线状态(用 dataclass,不用裸 dict)
├── sglang_model.py # 注册到 HF 架构串上的 SGLang 模型类
├── callbacks.py # 反馈式 AR 的三个回调,按需
├── routing.py # 数据驱动的动态路由,按需
└── components/ # 模型模块、预处理器、声码器、适配器

注册是自动发现的registry.py 遍历 models/ 下每个子包,收集 EntryClass,用它的 architecture 属性去匹配模型的 HF 配置。没有任何需要手工维护的清单。

真正需要手工连线的只剩两处:

  • 把 SGLang 模型类插进 model_runner/sglang_model_runner.py::_register_omni_model。HF 架构串和类名对不上时,用 model_arch_override 传同一个 key。
  • 上游权重加载不干净时写一个 weight_loader.py。多数模型不需要。

二、流水线的最小形状

文件放哪定了,接下来定形状:这个模型要拆成几个阶段。

绝大多数 TTS 模型是三段 —— 准备、生成码本、还原波形。这也是框架推荐的最小形状,Qwen3-TTS、Voxtral-TTS、Fish S2-Pro 都是它。两种常见的变体是往里插一个阶段或者加一条流式边:

基础形状 · 三阶段Qwen3-TTS · Voxtral-TTS · Fish S2-Pro 都是这个SIMPLE SCHEDULERpreprocessing校验请求、取参考并分词、建提示词状态OMNI SCHEDULERtts_engine自回归生成音频或 codec tokenSIMPLE 或流式调度器vocoder码本还原成波形,终点阶段变体一 · 插一个编码器阶段preprocessingaudio_encoder在 AR 卡上跑一次tts_engine什么时候用:参考编码器重到值得单独占一次调度。Higgs TTS 的多码本参考嵌入就是这么处理的。变体二 · 引擎到声码器开流式边tts_enginestream_to=["vocoder"]vocodercan_accept_stream_before_payload=True什么时候用:音频要在整段生成完成之前就出去。S2-Pro 是这条路的参考实现。三阶段是下限不是简化:预处理是 CPU 密集、引擎要 KV cache、声码器要么攒批要么流式,三种调度需求塞进一个循环只能取最保守的那套。阶段顺序、终点标记、GPU 放置、扇出全部在 config.py 里声明式地写完,然后模块级暴露 EntryClass,并给类设上 architecture。剩下的注册、拓扑推导、进程编排、传输选择都是框架的事,接入者不碰。
AR 阶段的工厂函数用 build_sglang_server_argscreate_sglang_infrastructure 两个共享助手,拿到 SGLang 的模型 worker、树缓存、请求池、KV 分配器、模型配置五件套,包一层 OmniScheduler 就跑起来了 —— KV cache、批选择、抢占、请求上限全部是上游原件。

三、请求边界上的两个陷阱

陷阱一:端点默认值会静默盖掉模型默认值。 HTTP 层会填一套采样默认值(目前用的是 S2-Pro 那套)。对任何别的模型来说,这些值看起来跟用户显式指定的一模一样。解法是让请求带一个 explicit_generation_params 之类的列表,请求构建器据此区分「用户设的」和「端点填的」。凡是端点有主观意见的字段,都要这一套。

陷阱二:输入是异质的。 TTS 客户端会用好几个名字传文本(inputtext、有时是聊天结构),参考音频也有好几种形状(ref_audioref_text,或者一个 references[] 列表)。在请求构建器里做归一化和必填校验,别把这套逻辑漏到 AR 阶段 —— 一个坏请求应该在碰到 GPU 之前就失败。

文件写好、形状定了,接下来是请求从 HTTP 进来变成模型输入的那一小段。这一段代码量最少,坑最多。

先看一个真会发生的情形:你的模型默认 temperature 是 0.7,但用户没填这个字段。请求到你手上时,temperature 已经是 0.9 了 —— HTTP 层替你填了默认值,而那套默认值是照着另一个模型调的。

你分不清这个 0.9 是「用户真的要 0.9」还是「HTTP 层随手填的」。分不清就没法决定该不该用自己的默认值。

另外,构建器交给调度器的应该是一个带类型的 dataclass,不是自由格式的 dict —— AR 阶段每一步都要读这些字段。

四、张量的设备纪律

接下来这条最容易在代码评审里被放过,因为它不影响功能,只影响性能,而且影响得很隐蔽。

典型写法是预处理阶段准备好参考张量之后,顺手写一个 .cpu() —— 「保险起见,免得后面设备对不上」。看起来无害。

把这个张量的实际路径画出来就知道代价了:

写了 .cpu() · 张量弹了两次预处理算出参考张量在显存里.cpu()显存 → 内存,一次同步跨进程传过去走的是内存那份AR 阶段拿到内存 → 显存,又拷回去每个请求付一次同步 + 两次拷贝留在原地 · 全程不下卡预处理算出参考张量在显存里跨进程传显存句柄(CUDA IPC)数据一个字节都没动AR 阶段直接用零多余拷贝这条边本来就该走 CUDA IPC「保险起见」那一句的真实代价:它把一条本可以零拷贝的边,变成了必须走内存的边 —— 传输选择是从张量在哪推导的。一次同步几十微秒看着不多,但它在每个请求上都发生,而且挡住了 CUDA IPC 这条最快的通路。规矩只有一条:张量在哪产生就留在哪,直到真的有人需要 CPU 字节。
这类问题在评审里很难被发现,因为 .cpu() 那一行看起来完全合理,代价发生在两个阶段之外的传输层。判断方法是反过来问:这个张量下一个真正读它的人在哪张卡上?如果还在 GPU 上,中间任何一次落地都是白费。

规矩只有一条:张量在哪产生就留在哪,直到真的有人需要 CPU 字节。

落到四条具体做法:

  • 预处理应当直接在 AR 的设备与 dtype 上创建提示词和参考张量,或者在交接前恰好归一化一次
  • 调度器的请求数据应该带设备张量,除非消费方明确是 CPU 端的,否则不要存 CPU 副本;
  • 不需要梯度就 detach,只在明确的所有权边界上做 dtype 转换(预处理输出、反馈缓冲写入、最终声码器或 HTTP 序列化);
  • 为了算一个稳定缓存键而做的 CPU 物化只产出元数据,不能替换掉 prefill/decode 用的那个设备张量。

还有一条前缀里的坑:前缀中混入了连续嵌入时(Higgs 和 Qwen3-TTS 都这样),radix 缓存键必须由嵌入内容派生。 两个恰好共享同一串占位符 token id 的不同提示词,会别名到同一段 KV 前缀 —— 表现就是一个用户的音频泄漏到另一个用户那里。这也是 09 篇那套内容哈希存在的理由。

五、中止清理的三条竞态路径

这是接入者最常漏的地方。任何以 request_id 为键、存在调度器之外的共享状态(预处理交给 AR 的张量暂存、会话句柄、参考缓存),都要在三条路径上被释放:

同一份共享状态,有三个不同的时刻可能被遗留。按时间轴看:

路径中止发生在为什么容易漏
1预处理交接之前最直观,也是唯一多数人会想到的一条
2已交到 AR 侧、构建器消费之前归属在两个阶段之间,容易两边都不管
3预处理算完时请求早已被中止结果对象无人接收,表现为缓慢的内存泄漏

解法是把同一个清理函数挂成每一个碰到这份共享状态的调度器abort_callback,并且让它幂等。

幂等不是可选项

这个函数按设计就会被同一个 request_id 调用多次。写成「第二次调用直接报错」是错的,写成「第二次调用悄悄什么都不做」才对。

六、错误处理的四条硬规定

这一节记录的是一个真实的翻车:同一个 OOM 在四个模型上产生了四种不同的对外行为。

修复之前,同一个 OOM 在四个模型上产生了四种不同的对外行为

模型对外表现评价
S2-Pro / VoxtralHTTP 500正确 —— 但只是因为它们没做任何事
Ming-OmniHTTP 200,波形为 None静默失败,客户端得自己判空
Qwen3-OmniHTTP 200,全零张量最坏 —— 下游看起来是段正常音频

四种行为没有一种是有意设计的,全是异常被宽泛的 except Exception 吞掉之后的副产物。

修复的四条规定:

#规定为什么要这一条
统一在 run_batch() 里捕获:标记请求失败 → outbox → 协调器 → 非流式返回 500给失败一条唯一的、可预测的出口
模型执行器里禁止写 except Exception用 lint 规则强制① 的捕获是路径局部的,写在别处的宽泛捕获照样溜过去
架构上禁止兜底路径:要么成功,要么把异常交给调度器,没有第三条路「返回一个看起来像成功的假结果」比直接失败更难查
CI 故障注入:注入 OOM,逐模型验证对外失败信号正确检测是 ② 的补充,不是替代
能用工具挡住的规矩就别指望人

② 特意写成 lint 规则而不是评审纪律。理由很实际:①的统一捕获只覆盖 run_batch() 这一条路径,而一个写在 add_request 或嵌入加载路径里的 except Exception 完全绕得过去 —— 靠人看是看不住的。

流式请求返回不了 500 —— 响应头早发出去了。约定是:成功的流以一个明确的完成哨兵帧结尾,失败的流在发哨兵之前直接断开连接。客户端判断失败的依据是「连接结束了但没收到哨兵」。这个契约必须写进客户端,否则一次中途 OOM 在客户端看起来就是「音频提前结束了」。

七、评审前要有的测试

全部是不需要 GPU 的单元测试,正好对应上面几节的规矩:

覆盖什么对应上文
请求边界:采样默认值的保留、必填输入的校验第三节
调度器请求数据:设备与 dtype 不变量第四节
中止清理:三条竞态路径第五节
阶段本地行为:声码器攒批或流式,看你选了哪个第二节

端到端质量另跑共享的 TTS 基准:

python -m benchmarks.eval.benchmark_tts_seedtts --help

报 WER/CER、样本数、吞吐、rtf_mean某个语种或切分上落后时,写一句可能的原因(采样配置、codec 或声码器版本、文本正则化、评测设置),而不是把数字丢在那里让它自己说话。

最后这条要求看着像客套,其实是很实的:一个没有解释的差数字,在下一个人手里只能从头查一遍。

下一篇07 - 性能瓶颈与优化案例