VibeVoice 流式推理完整链路:900 多行 inference 怎么读

2026-08-18

翻 VibeVoice 仓库的时候,最容易卡住的地方不是模型结构图,而是这个问题:调用方在 demo/web/app.py 里只写了一句 self.model.generate(...),另一个线程就开始从队列里往外取音频块了。中间那一层到底发生了什么?

答案全在 vibevoice/modular/modeling_vibevoice_streaming_inference.py 这一个文件里,篇幅九百多行。下面沿着它走一遍。文件路径与函数名都以仓库当前代码为准,这个项目在持续更新,读的时候请对照仓库最新内容。

先看被堵死的那个入口

一般 Hugging Face 风格的模型,读起来的第一步是找 forward。这个文件里的 forward 一进去就抛异常:

raise RuntimeError(
    "Unified forward is disabled. Use `forward_lm`, `forward_tts_lm`, or `generate` instead."
)

docstring 里仓库代码自述了原因:推理流程是分阶段的(先走基础文本 LM,再走 TTS LM,加上 streaming 与 diffusion 在 generate 里处理),一个整体的调用会把必须的时序(prefill、窗口推进、语音 diffusion 采样)藏起来。

所以真正的入口只有三个:forward_lm(基础文本 LM 的一步)、forward_tts_lm(TTS LM 的一步,需要传入 LM 的 hidden states)、generate(完整流程)。读代码时把这三个当成三条独立的路径,比一开始就找统一入口省事得多。

进主循环之前:四套状态,不是一套

generate 开头有一段很密的准备代码,核心是 _build_generate_config_model_kwargs 被调了四次,分别产出:

  • 正样本的基础 LM(model_kwargs / input_ids
  • 正样本的 TTS LM(tts_lm_model_kwargs / tts_lm_input_ids
  • 负样本的基础 LM(negative_model_kwargs
  • 负样本的 TTS LM(tts_lm_negative_model_kwargs

负样本那两套的 input_ids 是用 tokenizer.convert_tokens_to_ids("<|image_pad|>") 填出来的一列常量。它存在的意义在后面 CFG 那一步才看得出来。

四套状态的初值不是现算的,而是从 all_prefilled_outputs 里取的,四个键逐字是 lmtts_lmneg_lmneg_tts_lmdemo/realtime_model_inference_from_file.py 里这个字典是 torch.load(voice_sample, ...) 直接读进来的——也就是说 voice preset 文件装的是预填好的 prompt 输出与 kv cache。对应地,vibevoice/processor/vibevoice_streaming_processor.pyprocess_input_with_cached_promptcached_prompt['lm']['last_hidden_state'].size(1) 当长度,造出一串 pad_id 充数的 input_ids(代码注释里写的是 pseudo input ids),真正有内容的是 tts_text_ids,即整段脚本的 token。同一个文件里的 __call__ 被显式改成抛 NotImplementedError,提示你只能走 process_input_with_cached_prompt 这条路。

还有一行断言值得先记住:

assert batch_size == 1, "Currently only supports batch size == 1"

generate 的 docstring 也逐字写了 “The function only supports batch size = 1 currently”。所有 batch 相关的循环写法在当前代码里都只会走一遍。

主循环:文本窗和语音窗交替推进

文件顶部有两个模块级常量:

TTS_TEXT_WINDOW_SIZE = 5
TTS_SPEECH_WINDOW_SIZE = 6

这是仓库当前代码里的取值,随版本可能变动。while True 的每一轮做两件事:先按 tts_text_window_indextts_text_ids 上切一个文本窗喂进去,再进一个定长的内层循环采样语音 latent。

切窗口这两行是理解全局的关键:

cur_input_tts_text_ids = tts_text_ids[:, tts_text_window_index*TTS_TEXT_WINDOW_SIZE:(tts_text_window_index+1)*TTS_TEXT_WINDOW_SIZE]
next_text_window_size = tts_text_ids[:, (tts_text_window_index+1)*TTS_TEXT_WINDOW_SIZE:(tts_text_window_index+2)*TTS_TEXT_WINDOW_SIZE].shape[1]

注意第二行:它不用当前窗口,而是提前看了下一个窗口的长度。文本切完就走 forward_lm,把得到的 outputs.last_hidden_state 作为 lm_last_hidden_state 传给 forward_tts_lm,同时把 tts_text_masks 置成全 1。forward_tts_lm 的 docstring 写明这个 mask 的语义是 1=text、0=speech,它会经过 self.model.tts_input_types 变成一个类型 embedding 加到输入上。

文本窗为空(脚本读完了)时,外层这一段整个跳过,内层的语音窗还继续跑——这就是脚本读完之后语音还能继续吐的原因。

状态维护:两个同名函数,差别在 num_new_tokens

这个文件里有两个 _update_model_kwargs_for_generation:一个是模块级函数,一个是类方法(override 了父类的)。读代码时很容易看串。

模块级那个接受 num_new_tokens,一次给 attention_mask 补上对应个数的 1,并把 cache_position 重新排成 torch.arange(cache_pos[-1] + 1, cache_pos[-1] + num_new_tokens + 1);类方法那个走父类实现,然后补一道 _ensure_cache_has_layers。后者是为 transformers 4.57 之后 cache 系统重构准备的兼容层,文件里用 MockCacheLayer 伪造出新版本期待的 layers 接口——文件顶部注释逐字写了 “Transformers >= 4.57 Compatibility Layer”。

真正容易读漏的是内层循环末尾这个分支:

if cur_speech_index == TTS_SPEECH_WINDOW_SIZE - 1 and next_text_window_size > 0:
    tts_lm_model_kwargs = _update_model_kwargs_for_generation(
        tts_lm_outputs, tts_lm_model_kwargs, num_new_tokens=next_text_window_size,
    )
else:
    tts_lm_model_kwargs = self._update_model_kwargs_for_generation(...)

只有在语音窗的最后一步、且下一个文本窗非空时,才用多 token 版本,一次把下一段文本要占的位置在 attention_mask 与 cache_position 上腾出来。其余每一步都是单 token 推进。之前提前算 next_text_window_size 就是为了这里。

一步语音是怎么生成又喂回去的

内层循环里,一轮的顺序是这样的:

  1. tts_lm_outputs.last_hidden_statetts_lm_negative_outputs.last_hidden_state 各取最后一个位置,作为 positive_conditionnegative_condition
  2. sample_speech_tokens。这个函数把正负条件在 batch 维拼起来,从 torch.randn(..., self.config.acoustic_vae_dim) 出发,按 self.model.noise_scheduler.timesteps 逐步去噪;每一步用 uncond_eps + cfg_scale * (cond_eps - uncond_eps) 合成,最后只返回前一半。scheduler 来自 vibevoice/schedule/dpm_solver.pyDPMSolverMultistepScheduler,步数由 set_ddpm_inference_steps 设置(不传就回落到 config.diffusion_head_config.ddpm_num_inference_steps)。
  3. 反归一化:speech_latent / self.model.speech_scaling_factor - self.model.speech_bias_factor,然后交给 self.model.acoustic_tokenizer.decode,带上 VibeVoiceTokenizerStreamingCache 实例并置 use_cache=True。这个 cache 在整个 generate 里只建一次;它的类 docstring 自述是「流式卷积的 cache,类似 attention 里的 KV cache」。
  4. 把同一个 latent 过 self.model.acoustic_connector 得到 acoustic_embed,作为下一步 forward_tts_lmlm_last_hidden_statetts_text_masks 这次置 0。

第 4 步是闭环所在:刚生成的声学 latent 换个投影就成了下一步的输入。语音 token 位置上的 tts_lm_input_ids 只是拿 torch.ones_like 拼了个占位 id,真正承载信息的是被覆盖进去的 embedding。

什么时候停

循环有四个出口,都在代码里写死:外部传入的 stop_check_fn() 返回真、finished_tags.all()tts_lm_input_ids.shape[1] 超过 max_length(文本窗与语音窗后各判一次),以及 EOS:

tts_eos_logits = torch.sigmoid(self.tts_eos_classifier(tts_lm_outputs.last_hidden_state[diffusion_indices, -1, :]))
if tts_eos_logits[0].item() > 0.5:

这里有一处值得把两段代码放一起看:_build_generate_config_model_kwargsreturn_processors=True 时构造了 logits_processorstopping_criteria 并返回,但在 generate 的主循环里,这两个变量此后没有再被引用;停止判定实际由上面那个二分类器、max_lengthstop_check_fn 承担。至于为什么这么写,仓库里我们没有找到对应说明,这里只陈述代码现状。

另外两处默认值也可以并排看:generate 签名里 cfg_scale: float = 1.0sample_speech_tokens 签名里 cfg_scale=3.0,而 demo/realtime_model_inference_from_file.py--cfg_scale 参数默认 1.5。实际生效的是调用方传进 generate 的那个值。这些都是仓库当前代码里的默认值,随版本可能变动。

还有一个:demo/web/app.pygenerate 时传了 refresh_negative=...,但模型侧代码里我们没有搜到这个参数名,它会落进 **kwargs。读 demo 时别把它当成 generate 的正式参数。

与 streamer 的交接点只有三行

整个 generate 里碰 audio_streamer 的地方就三处:语音块生成后 audio_streamer.put(audio_chunk, diffusion_indices);某个样本判定 EOS 后 audio_streamer.end(diffusion_indices);以及主循环退出后无条件的 audio_streamer.end()。此外 stop_check_fn 触发时也会先 end() 再 break。

vibevoice/modular/streamer.pyAudioStreamer 就是每个样本一个 Queue,加一组 finished_flagsput 会跳过已结束的样本,end 往队列里塞 stop_signalAsyncAudioStreamer 继承它,把队列换成 asyncio.Queue,构造时记下 asyncio.get_running_loop(),并在 put / end 里用 self.loop.call_soon_threadsafe 把值投递回那个 loop。仓库里没有写明这么设计的理由,只能并排看一个事实:demo/web/app.py 确实是把 generate 放进 threading.Thread 里跑的。

调用方那侧的完整姿势在 demo/web/app.py 里:AudioStreamer(batch_size=1, stop_signal=None, timeout=None),把 generate 放进 threading.Thread,主线程 audio_streamer.get_stream(0) 迭代取块,stop_check_fn=stop_event.is_setfinally 里依次 stop_signal.set()audio_streamer.end()thread.join()。这三步顺序值得照抄——只 set 事件不 end 队列,取流的一侧可能还在等。

顺带一个容易忽略的事实:即使传了 audio_streamergenerate 内部仍然在往 audio_chunks 里累积,结束时 torch.cat 成完整音频返回;return_speech=False 时才把这段拼接跳过(返回 speech_outputs=None)。

边界

  • batch_size == 1 是断言,不是建议。
  • forward 被设计成必定抛错,set_output_embeddings 同样直接抛 RuntimeErrorget_output_embeddings 返回 None(docstring 写明这个模型没有 lm_head)。
  • docs/vibevoice-realtime-0.5b.md 的 TODO 列表里,“Implement streaming text input function to feed new tokens while audio is still being generated” 这一项还没有打勾;同一份文档也写明该实时变体只支持单说话人、主要面向英语,其它语言可能产生不可预测的结果。代码里有窗口化的文本喂入机制,不等于流式文本输入这个功能已经完成。
  • 设备分支:demo/realtime_model_inference_from_file.py--device 默认值是 cudampscpu 的三级表达式。代码里还专门把 mpx 这个拼写纠正成 mps,并且在指定了 mps 却不可用时打印警告回落到 CPU。走 mps 那条分支时仓库把 dtype 固定成 torch.float32、注意力实现固定成 sdpa,注释逐字写着 flash_attention_2 在 MPS 上不受支持。Windows 侧只会落到 cudacpu 这两个分支——需要说明的是,「MPS 是 PyTorch 在 Apple 平台上的后端」这句是通用背景,不是本仓库写的内容。另外 flash_attention_2 加载失败时会回退到 sdpa,回退时打印的原文提示只有 flash_attention_2 经过完整测试。

以上代码片段原样引自仓库文件,未经实测,以仓库最新代码为准。

读完这条链路再回头看模型卡和架构图会顺很多:窗口常量决定节奏,四套 kv cache 决定状态怎么走,acoustic_connector 决定闭环怎么合上,而 streamer 只是最后那三行的搬运工。


本文依据 github.com/microsoft/VibeVoice 仓库与 Hugging Face 模型卡于 2026-08-18 的公开内容整理, 事实来自仓库内的文档与源码。我们没有下载权重、没有跑过推理、也没有做过训练, 因此不涉及显存占用、推理速度、识别准确率与音质的任何描述,也不与其它模型做比较或排名。 该项目持续更新,文中涉及的模块路径、配置字段与接口写法随版本变动,请以仓库最新内容为准。

需要说明的是:仓库 README 记载,2025-09-05 微软因发现有与既定意图不符的使用方式, 基于负责任 AI 原则从该仓库移除了 VibeVoice-TTS 代码; 当前 vibevoice/modular/modeling_vibevoice.py 首行注释标明其来自社区 fork, 且该模块未被 vibevoice/modular/__init__.py__all__ 导出。 本文只讲代码与架构,不构成 TTS 推理的可用性保证。

仓库 README 的风险与限制一节写明:该模型仅供研究与开发用途, 未经进一步测试与开发不建议用于商业或真实场景,并特别提示了合成语音被用于伪造与虚假信息的风险。 使用合成语音时应遵守所在司法辖区的法律法规,并在分享 AI 生成内容时主动披露。

想系统学会用 AI?报名体系课或加入会员,照着学、照着用。