文本到语音与语音 文本到语音(Text-to-Speech,TTS)合成把 ASR 流水线逆转过来,从书面文字生成自然的音频。本文件涵盖 TTS 流水线(文本归一化、G2P、声学模型、声码器)、Tacotron、WaveNet、HiFi-GAN、声音克隆、语音转换以及语音活动检测(VAD)。 在第 1 个文件中我们构建了信号处理工具箱:波形、频谱图、梅尔滤波器组和 MFCC。在第 2 个文件中我们把语音转换成了文字。现在我们把箭头反过来:给定文字,合成自然的语音。这就是文本到语音(TTS),这个问题同时也打开了通往语音转换、声音克隆和语音活动检测的大门。 可以把 TTS 想象成一场舞台表演。剧本就是文本输入。导演(声学模型)决定每一句台词该怎么听起来——音高、节奏、重音。
文本到语音(Text-to-Speech,TTS)合成把 ASR 流水线逆转过来,从书面文字生成自然的音频。本文件涵盖 TTS 流水线(文本归一化、G2P、声学模型、声码器)、Tacotron、WaveNet、HiFi-GAN、声音克隆、语音转换以及语音活动检测(VAD)。
在第 1 个文件中我们构建了信号处理工具箱:波形、频谱图、梅尔滤波器组和 MFCC。在第 2 个文件中我们把语音转换成了文字。现在我们把箭头反过来:给定文字,合成自然的语音。这就是文本到语音(TTS),这个问题同时也打开了通往语音转换、声音克隆和语音活动检测的大门。
可以把 TTS 想象成一场舞台表演。剧本就是文本输入。导演(声学模型)决定每一句台词该怎么听起来——音高、节奏、重音。乐队(声码器)随后演奏总谱,产生观众听到的实际声波。现代神经 TTS 用堪比真人说话的表演取代了基于规则的系统那种僵硬、机械的念白。
文本到语音流水线 标准 TTS 流水线有四个阶段:(1) 文本归一化,(2) 音素转换,(3) 声学模型,(4) 声码器。一些现代系统把阶段 3 和 4 合并成一个单一的端到端模型,但概念上的分解仍然有用。
文本归一化把原始文本转换为可发音的形式。缩写被展开("Dr." 变为 "Doctor"),数字变为词("1984" 变为 "nineteen eighty-four"),货币符号被读出("$5" 变为 "five dollars"),URL 或特殊字符被处理。这一阶段通常是基于规则的,配以特定语言的语法,不过也存在神经归一化模型。这里的错误会传播到下游每个阶段:如果 "St." 被读成 "saint" 而不是 "street",整句都会错。
字素到音素(Grapheme-to-Phoneme,G2P)转换把归一化的文本映射到音素序列。英语的拼写出了名地不规则("though"、"through"、"tough" 中 "ough" 的发音各不相同),所以词典查找(CMU 发音词典)处理常见词,而神经序列到序列模型(第 6 章的编码器-解码器或第 7 章的 Transformer)处理词表外的词。浅层正字法的语言(西班牙语、芬兰语)需要更简单的 G2P。输出通常是国际音标(IPA)序列或等价的内部音素集。
声学模型接收音素序列并产生中间的声学表示,几乎总是梅尔频谱图(第 1 个文件)。梅尔频谱图捕捉每一时间帧的谱包络,编码了声码器重建波形所需的感知相关信息。声学模型必须决定时长(每个音素持续多久)、音高(基频 F_0)和能量(响度)。
声码器接收梅尔频谱图并产生原始音频波形。这是一个不适定的反演问题:许多波形都能产生相同的频谱图,因为相位信息被丢弃了。经典声码器(Griffin-Lim、WORLD)使用迭代或信号模型方法,但神经声码器在质量上现在已经占主导。
声码器:WaveNet(van den Oord 等,2016)是第一个能产生几乎与人类录音难以区分的语音的神经声码器。它以自回归方式对波形建模,在所有之前的样本条件下预测每个样本 x_t:
其中 c 是条件信号(梅尔频谱图)。每个样本是 16 位的,所以对 65536 个值做朴素的 softmax 是不现实的。WaveNet 使用 mu-law 压扩把量化电平减少到 256 个,后来的变体使用逻辑斯谛分布的混合。
WaveNet 的核心构建块是膨胀因果卷积。因果意味着滤波器权重只看过去的样本(没有未来泄露)。膨胀意味着滤波器以指数级增大的间隔跳过样本:膨胀因子 1, 2, 4, 8, \ldots, 512。这给出了指数级大的感受野,同时保持参数量为线性。
每层的门控激活为:
其中 W_f 和 W_g 是滤波器和门卷积权重,\ast 表示膨胀因果卷积,\odot 是逐元素相乘。这种门控机制(来自第 6 章的 LSTM)让网络能够控制信息流。
WaveNet 产生了卓越的质量,但在推理时却慢得令人痛苦:生成一秒 24 kHz 音频需要 24000 次顺序前向传播。这推动了随后的所有声码器研究。
WaveRNN(Kalchbrenner 等,2018)用单层循环网络取代了 WaveNet 的深层卷积堆栈。它把每个 16 位样本分成粗(高 8 位)和细(低 8 位)两部分,用一个 GRU(第 6 章)分别预测。这种双 softmax 方法在保持高质量的同时显著减少了计算。通过精心的内核优化,WaveRNN 在移动 CPU 上足以实时运行。
WaveGlow(Prenger 等,2019)是一种基于流的声码器,完全避免了自回归生成。它使用一系列可逆变换(仿射耦合层,第 6 章的归一化流)把一个简单的高斯分布映射到波形分布。训练使用变量替换公式最大化精确对数似然:
其中 z = f(x) 是把 x 通过流传递得到的潜变量。在推理时,从 z \sim \mathcal{N}(0, I) 采样,并在一次并行传递中通过逆流推送。WaveGlow 用模型大小(耦合层用大网络)换取生成速度。
HiFi-GAN(Kong 等,2020)使用生成对抗网络从梅尔频谱图合成波形。生成器通过一系列转置卷积对梅尔频谱图上采样,每个后接一个多感受野融合(Multi-Receptive Field Fusion,MRF)模块。MRF 模块并行地应用多个具有不同核大小和膨胀率的残差块,然后对它们的输出求和。这使得生成器能够同时捕捉多个时间尺度的模式。
HiFi-GAN 使用两种判别器。多周期判别器(Multi-Period Discriminator,MPD)把一维波形按不同周期(2、3、5、7、11)折叠成二维,然后施加二维卷积。这捕捉了不同基频处的周期结构。多尺度判别器(Multi-Scale Discriminator,MSD)分别在原始波形、2 倍下采样和 4 倍下采样版本上操作,捕捉不同时间分辨率的模式。
训练目标结合了对抗损失、梅尔频谱图重建损失(合成音频与真实音频梅尔频谱图之间的 L1 距离)和特征匹配损失(判别器中间特征之间的 L1 距离):
HiFi-GAN 实现了可与 WaveNet 媲美的合成质量,同时快 1000 倍以上,使在单张 GPU 上实时生成成为可能。
神经源-滤波器(Neural Source-Filter,NSF)模型把传统信号处理与神经网络结合。在经典的源-滤波器模型中,浊音由一个源激励(基频 F_0 处的周期脉冲串)通过一个声道滤波器(谱包络)产生。NSF 模型用神经网络取代手工设计的滤波器,同时保留显式的源信号。输入的 F_0 轮廓提供了纯数据驱动的声码器有时难以应对的精细音高控制。
声学模型:Tacotron(Wang 等,2017)是第一个直接把字符序列转换为梅尔频谱图的端到端神经 TTS 系统。它使用带注意力的编码器-解码器架构(第 7 章)。编码器用卷积库、高速公路网络和双向 GRU 处理字符/音素序列。解码器是一个自回归 GRU,一次预测一帧梅尔,使用前一帧和注意力上下文作为输入。
Tacotron 2(Shen 等,2018)显著改进了架构。编码器是一个 3 层一维卷积堆栈,后接一个双向 LSTM(第 6 章)。解码器是一个带位置敏感注意力的 2 层 LSTM,它不仅把注意力机制条件化于编码器输出和解码器状态,还条件化于之前步骤的累积注意力权重。这防止了注意力跳过或重复词语这种常见的失败模式。
其中 s_{i-1} 是前一个解码器状态,h_j 是位置 j 处的编码器输出,f_{i,j} 是通过把累积注意力权重 \sum_{k<i} \alpha_{k,j} 与一维卷积滤波器做卷积得到的位置特征。注意力权重为 \alpha_{i,j} = \text{softmax}(e_{i,j})。
Tacotron 2 的解码器还在每步预测一个停止 token 概率,指示梅尔频谱图何时完成。输出的梅尔频谱图随后被传递给声码器(最初是 WaveNet,后来被 HiFi-GAN 或类似物取代)。
Tacotron 2 的自回归特性意味着合成速度受限于梅尔帧的数量。对于一个典型的每秒 80 帧的梅尔频谱图,5 秒的话语需要 400 个顺序解码步。
FastSpeech(Ren 等,2019)用非自回归声学模型解决了速度问题。FastSpeech 不再顺序地生成梅尔帧,而是并行生成所有帧。关键挑战是确定每个音素应该产生多少梅尔帧,FastSpeech 用一个时长预测器来处理。
时长预测器是一个小型卷积网络,预测每个音素的整数时长(梅尔帧数)。训练期间,从预训练的自回归教师模型(Tacotron 2)的注意力对齐中提取真实时长。推理期间,使用预测的时长通过一个长度调节器把音素级隐藏序列扩展到帧级,该调节器简单地把每个音素的隐藏表示重复预测的帧数。
FastSpeech 2(Ren 等,2021)通过去除教师-学生蒸馏改进了 FastSpeech。它使用强制对齐(来自第 2 个文件的声学模型框架)直接提取真实时长,并除了时长外还添加了显式的音高(F_0)和能量方差适配器。每个适配器都是一个小型卷积预测器,其输出条件化解码器:
\begin{aligned}
\hat{d}_i &= \text{DurationPredictor}(h_i) \\
\hat{p}_i &= \text{PitchPredictor}(h_i) \\
\hat{e}_i &= \text{EnergyPredictor}(h_i)
\end{aligned}
其中 h_i 是音素 i 的编码器隐藏状态。训练时使用真实值;推理时,预测的值给予对韵律的显式控制。这种可控性是 FastSpeech 2 的主要优势:调整音高、速度或能量就像缩放预测器输出一样简单。
FastSpeech 2 在推理时通常比 Tacotron 2 快 10-20 倍,并避免了常见的自回归失败模式,如跳词、重复和注意力坍塌。
VITS(Kim 等,2021)是一个端到端 TTS 模型,直接从文本生成波形,消除了独立的声码器阶段。VITS 把条件变分自编码器(第 6 章)与归一化流和对抗训练结合。后验编码器把真实梅尔频谱图映射到潜空间,先验编码器把音素(通过基于 Transformer 的文本编码器和时长预测器)映射到同一潜空间,解码器(基于 HiFi-GAN)从潜样本生成波形。
VITS 的训练目标结合了:
VITS 产生了比两阶段系统(FastSpeech 2 + HiFi-GAN)更高的质量,因为声学模型和声码器被联合优化,避免了预测梅尔频谱图与真实梅尔频谱图之间的不匹配(这种不匹配会降低两阶段系统的质量)。
VALL-E(Wang 等,2023)从根本上把 TTS 重新构架为离散音频 token 上的语言建模问题。它使用神经音频编解码器(EnCodec)把语音表示为来自多个码本层级的离散码序列。给定一个文本提示和一段 3 秒的注册话语(也编码为离散 token),VALL-E 使用 Transformer 语言模型自回归地预测音频 token。
VALL-E 使用两个模型:一个自回归(AR)模型逐个 token 地生成第一个码本层级,以及一个非自回归(NAR)模型在第一个层级和彼此条件下并行预测其余码本层级。这种编解码器语言模型方法实现了卓越的零样本声音克隆:3 秒的样本足以复现一个说话人的声音、音色甚至情感语调。
StyleTTS(Li 等,2022)和 StyleTTS 2 把语音解耦为内容和风格组件。风格编码器从参考音频中提取风格向量,捕捉说话人身份、韵律和录音条件。推理时,风格可以从学习到的先验分布中采样,也可以从参考话语中迁移。StyleTTS 2 使用扩散模型(第 8 章)作为风格先验,生成多样且自然的韵律。
Kokoro(2024)是一个轻量、高质量的开源 TTS 模型,以其小巧的尺寸(约 82M 参数)和令人印象深刻的自然度而著称。它使用受 StyleTTS 2 启发的架构,配以基于扩散的风格先验和一个微调过的 ISTFTNet 声码器,该声码器直接预测 STFT 系数(来自第 1 个文件)而非原始波形样本。尽管体积只是 VALL-E 等模型的一小部分,Kokoro 在英语、日语、法语、韩语和中文上实现了接近人类的自然度,证明了精心策划的训练数据和高效的架构设计可以与蛮力规模竞争。Kokoro 的小巧使其适合本地和边缘部署。
Orpheus(Canopy Labs,2025)是一系列基于 VALL-E 所开创的编解码器语言模型范式的开源 TTS 模型(1B 和 3B 参数)。Orpheus 用一个 LLM 骨干(微调过的 Llama 3)进一步推进了这一思想,该骨干直接生成 SNAC 音频编解码器 token。它的突出特点是类人的情感表现力:它以非凡的自然度处理笑声、叹息、犹豫和情感韵律。Orpheus 可以通过输入文本中的 [laugh] 或 [sigh] 等标签来提示,对副语言表达给予细粒度控制。
Dia(Nari Labs,2025)是一个开源的对话 TTS 模型,能从单段文本转录生成逼真的多人对话。基于一个 1.6B 参数的编码器-解码器 Transformer,Dia 处理对话中的轮流发言、特定说话人的声音和非语言线索(笑声、停顿)。它还支持从短音频提示进行声音克隆,在对话上下文中实现零样本说话人生成。
Sesame CSM(Conversational Speech Model,对话语音模型,2025)专注于自然的多人轮替对话语音。Sesome 不优化朗读风格的 TTS,而是建模真实对话的动态:反向回应("uh huh")、打断、说话人之间的节奏变化和情感响应性。该模型使用以对话上下文(文本和音频历史都包括)为条件的 Transformer 骨干,产生能够根据对话流调整风格的语音。
Fish Speech(Fish Audio,2024)是一个开源 TTS 系统,使用双自回归架构:一个大型语言模型从文本生成语义 token,一个较小的模型把这些转换为 VQGAN 声学 token,再由声码器解码为波形。Fish Speech 支持从 10-15 秒参考进行零样本声音克隆,并实现了适合实时应用的低延迟。其模块化设计允许独立地替换组件(例如不同的声码器)。
ChatTTS(2024)是一个开源的对话 TTS 模型,专为聊天机器人和虚拟助手等对话应用设计。它生成自然、对话风格的语音,并通过嵌入在文本输入中的特殊 token 对韵律特征(笑声、停顿、填充词)进行细粒度控制。ChatTTS 支持中英混合合成和多人语音生成。
Bark(Suno,2023)是一个基于 Transformer 的开源模型,能从文本提示生成语音、音乐和音效。它使用一个三阶段的 Transformer 模型流水线(文本 → 语义 token → 粗声学 token → 精细声学 token),支持声音克隆、多语言合成以及非语音音频(如音乐和环境声)。Bark 的通用性以可控性为代价——它不如专用 TTS 系统精确,但更灵活。
Parler-TTS(Hugging Face,2024)采用自然语言描述方法来控制声音:用户无需参考音频片段来指定风格,而是提供一段文本描述,如"一位声音温暖、富有表现力的女性说话人在安静的房间里说话"。Parler-TTS 在标注的语音数据上训练,每个话语都配有说话风格的自然语言描述,从而无需任何参考音频即可直观控制。
Neuphonic 是一个基于 API 的 TTS 平台,针对超低延迟语音合成进行了优化,面向实时语音代理和对话式 AI 应用。它通过一种流式架构实现低于 100 ms 的首音频时间,该架构在完整输入文本可用之前就开始生成音频。Neuphonic 专注于部署和延迟优化层,而非新颖的模型架构,围绕现代神经 TTS 提供生产级基础设施。
KittenTTS 是一个为效率和低资源部署设计的紧凑、快速的 TTS 模型。它优先考虑最小的延迟和小的模型尺寸,用于边缘和嵌入式应用,以一定的自然度换取在 CPU 和移动设备上的实时性能。
现代 TTS 的格局正在分化为两种范式:(1) 编解码器语言模型(VALL-E、Orpheus、Fish Speech),把语音生成视为离散音频码上的下一个 token 预测,利用 LLM 的缩放定律;(2) 基于流/扩散的模型(VITS、StyleTTS 2、Kokoro),通过迭代细化生成连续的梅尔频谱图或波形。编解码器语言模型在零样本克隆和表现力方面表现出色;流/扩散模型往往更小更快。两者都在快速向人类水平的自然度收敛。
韵律建模控制语音的"音乐性":音高、时长、能量、节奏和语调。没有好的韵律,即使单个音素清晰,合成的语音听起来也很平淡、机械。可以把韵律想象成单调的 GPS 语音与富有表现力的有声书朗读者之间的差别。
音高(基频 F_0)是语音感知到的高低。它在问句结尾升高,在陈述句结尾下降,并在情感语音中持续变化。F_0 通过 CREPE(一种神经音高跟踪器)或 YIN(基于自相关,来自第 1 个文件)等算法从音频中提取。在 TTS 中,音高要么由声学模型预测(FastSpeech 2 的音高预测器),要么被隐式学习(Tacotron 2)。
时长决定语速和节奏。重读音节更长,功能词被缩短,停顿标记短语边界。时长建模在非自回归模型中是显式的(FastSpeech),在自回归模型中是隐式的(Tacotron 的注意力对齐决定时长)。
能量(响度)承载强调。"I didn't say HE stole it" 与 "I didn't say he STOLE it" 通过能量模式传达完全不同的含义。
风格嵌入捕捉更高层次的韵律模式。全局风格 token(Global Style Token,GST)框架(Wang 等,2018)学习一银行风格 token(对一个学习到的嵌入集做软注意力),这些 token 捕捉诸如"兴奋"、"悲伤"或"耳语"等说话风格。风格嵌入从参考话语中提取并加到编码器输出上,允许在推理时进行风格迁移。
语音转换(Voice Conversion,VC)在保留语言内容的同时改变话语的说话人身份。想象一下录制你自己的声音,然后让输出听起来像某个特定的目标说话人。VC 需要把说话人身份与内容解耦。
说话人嵌入(在第 4 个文件中详述)把说话人身份编码为固定维向量。这些可以来自预训练的说话人确认模型(x-vector、ECAPA-TDNN)。在 VC 中,源语音被编码为与说话人无关的内容表示,然后用目标说话人嵌入解码。
解耦表示把语音分离为独立因子:内容(音素)、说话人身份、音高和节奏。方法包括:
声音克隆用目标说话人的声音合成语音。多人 TTS 在许多说话人的数据上训练,把模型条件化于说话人嵌入。推理时,从注册音频中提取新说话人的嵌入,并用于条件化生成。
少样本声音克隆使用少量数据(几分钟)适应新说话人。说话人编码器从注册音频中提取嵌入,TTS 模型在该嵌入条件下生成语音。这是 SV2TTS(Jia 等,2018)使用的方法:一个单独训练的说话人编码器、一个以说话人嵌入为条件的 Tacotron 2 合成器,以及一个 WaveRNN 声码器。
零样本声音克隆完全不需要适应:一段简短的话语(3-30 秒)就足够了。VALL-E 通过把注册音频作为语言模型的提示来实现这一点。模型学会以同一声音继续生成,因为它在大规模多人数据上训练,其中话语内的声音一致性是统计常态。
语音活动检测(Voice Activity Detection,VAD)在每个时间帧回答一个简单的二元问题:是否有人在说话?尽管简单,VAD 是 ASR(第 2 个文件)、说话人分离(第 4 个文件)和降噪(第 5 个文件)的关键预处理步骤。好的 VAD 通过跳过静音来减少计算,并通过防止噪声被当作语音处理来提高准确率。
经典 VAD 使用能量阈值(语音比静音响)、过零率(语音有特征的过零模式)和频谱特征。这些在低信噪比的嘈杂环境中会失效。
神经 VAD 模型把问题作为帧级二元分类。一个小型 RNN 或 CNN 接收声学特征(来自第 1 个文件的对数梅尔能量)并预测语音/非语音概率。
WebRTC VAD(Google)是一个经典的轻量级 VAD,使用基于 GMM 的分类器作用于简单频谱特征。它在四个激进级别(0-3)上操作,非常快,但在音乐、非语音发声和低信噪比环境中表现不佳。由于其零依赖的简洁性,它仍被广泛用作基线。
Silero VAD(Silero Team,2021)是事实上的生产用神经 VAD 标准。其架构是一小堆深度可分离的一维卷积(第 8 章的 MobileNet 思想应用于音频),后接一个用于时间上下文的单层 LSTM,最后是一个线性头,每帧产生一个语音概率。整个模型不到 2MB(约 1M 参数),以 30-100 ms 的块处理音频。
声学活动检测(Acoustic Activity Detection,AAD)把 VAD 推广到检测任何声学活动,不只是语音。这在智能家居设备、安防系统和野生动物监测中很有用。AAD 模型检测诸如玻璃破碎、狗叫或警报等事件,通常使用第 4 个文件中描述的音频分类框架。
TTS 的评估指标衡量客观质量和主观自然度:
import jax import jax.numpy as jnp import matplotlib.pyplot as plt # 生成合成波形(谐波之和,模拟元音) sr = 16000 duration = 1.0 t = jnp.linspace(0, duration, int(sr * duration)) f0 = 220.0 # 基频 waveform = ( 0.6 * jnp.sin(2 * jnp.pi * f0 * t) + 0.3 * jnp.sin(2 * jnp.pi * 2 * f0 * t) + 0.1 * jnp.sin(2 * jnp.pi * 3 * f0 * t) ) # 计算 STFT n_fft = 1024 hop_length = 256 window = jnp.hanning(n_fft) def stft(signal, n_fft, hop_length, window): """计算短时傅里叶变换。""" n_frames = 1 + (len(signal) - n_fft) // hop_length frames = jnp.stack([ signal[i * hop_length : i * hop_length + n_fft] * window for i in range(n_frames) ]) return jnp.fft.rfft(frames, n=n_fft) def istft(stft_matrix, hop_length, window, length): """用重叠相加计算逆 STFT。""" n_fft = (stft_matrix.shape[1] - 1) * 2 n_frames = stft_matrix.shape[0] frames = jnp.fft.irfft(stft_matrix, n=n_fft) frames = frames * window[None, :] output = jnp.zeros(length) for i in range(n_frames): start = i * hop_length end = start + n_fft if end <= length: output = output.at[start:end].add(frames[i]) return output # 正向 STFT S = stft(waveform, n_fft, hop_length, window) magnitude = jnp.abs(S) # 梅尔滤波器组 n_mels = 80 mel_low = 0.0 mel_high = 2595 * jnp.log10(1 + (sr / 2) / 700) mel_points = jnp.linspace(mel_low, mel_high, n_mels + 2) hz_points = 700 * (10 ** (mel_points / 2595) - 1) freq_bins = jnp.floor((n_fft + 1) * hz_points / sr).astype(int) mel_filterbank = jnp.zeros((n_mels, n_fft // 2 + 1)) for m in range(n_mels): f_left = freq_bins[m] f_center = freq_bins[m + 1] f_right = freq_bins[m + 2] for k in range(f_left, f_center): mel_filterbank = mel_filterbank.at[m, k].set( (k - f_left) / max(f_center - f_left, 1) ) for k in range(f_center, f_right): mel_filterbank = mel_filterbank.at[m, k].set( (f_right - k) / max(f_right - f_center, 1) ) # 到梅尔再回来(伪逆) mel_spec = magnitude @ mel_filterbank.T magnitude_reconstructed = mel_spec @ jnp.linalg.pinv(mel_filterbank.T) magnitude_reconstructed = jnp.maximum(magnitude_reconstructed, 1e-7) # Griffin-Lim 算法 def griffin_lim(magnitude, n_iter, hop_length, window, signal_length): """迭代相位重建。""" n_fft = (magnitude.shape[1] - 1) * 2 key = jax.random.PRNGKey(42) phase = jax.random.uniform(key, magnitude.shape, minval=-jnp.pi, maxval=jnp.pi) for _ in range(n_iter): complex_spec = magnitude * jnp.exp(1j * phase) signal = istft(complex_spec, hop_length, window, signal_length) reanalysis = stft(signal, n_fft, hop_length, window) phase = jnp.angle(reanalysis) complex_spec = magnitude * jnp.exp(1j * phase) return istft(complex_spec, hop_length, window, signal_length) reconstructed = griffin_lim(magnitude_reconstructed, n_iter=60, hop_length=hop_length, window=window, signal_length=len(waveform)) # 对比绘图 fig, axes = plt.subplots(3, 1, figsize=(12, 8)) axes[0].plot(t[:1000], waveform[:1000], color='#3498db', linewidth=0.8) axes[0].set_title('Original Waveform') axes[0].set_ylabel('Amplitude') axes[1].imshow(jnp.log1p(mel_spec.T), aspect='auto', origin='lower', cmap='magma') axes[1].set_title('Mel Spectrogram (intermediate representation)') axes[1].set_ylabel('Mel bin') axes[2].plot(t[:1000], reconstructed[:1000], color='#e74c3c', linewidth=0.8) axes[2].set_title('Griffin-Lim Reconstructed Waveform (60 iterations)') axes[2].set_xlabel('Time (s)') axes[2].set_ylabel('Amplitude') plt.tight_layout() plt.show() # 衡量重建误差 mse = jnp.mean((waveform[:len(reconstructed)] - reconstructed[:len(waveform)]) ** 2) print(f"MSE between original and reconstructed: {mse:.6f}") print("Note: phase information loss through mel inversion causes artifacts.")
import jax import jax.numpy as jnp import jax.random as jr import matplotlib.pyplot as plt # 模拟带有真实时长的音素序列 # 在真实 TTS 中,时长来自强制对齐或教师注意力 def generate_synthetic_data(key, n_samples=200, max_phonemes=30, embed_dim=64): """生成合成的音素嵌入和时长。""" keys = jr.split(key, 4) lengths = jr.randint(keys[0], (n_samples,), 5, max_phonemes) all_embeddings = [] all_durations = [] all_masks = [] for i in range(n_samples): L = int(lengths[i]) emb = jr.normal(keys[1], (max_phonemes, embed_dim)) # 时长:元音(偶数索引)更长,辅音更短 base_dur = jnp.where(jnp.arange(max_phonemes) % 2 == 0, 8.0, 4.0) noise = jr.normal(jr.fold_in(keys[2], i), (max_phonemes,)) * 1.5 dur = jnp.clip(base_dur + noise, 1.0, 20.0).astype(jnp.float32) mask = (jnp.arange(max_phonemes) < L).astype(jnp.float32) all_embeddings.append(emb) all_durations.append(dur * mask) all_masks.append(mask) return (jnp.stack(all_embeddings), jnp.stack(all_durations), jnp.stack(all_masks)) key = jr.PRNGKey(42) embeddings, durations, masks = generate_synthetic_data(key) # 时长预测器:2 层一维卷积 + 线性投影 def init_duration_predictor(key, embed_dim=64, hidden_dim=128, kernel_size=3): """初始化时长预测器权重。""" keys = jr.split(key, 4) scale1 = jnp.sqrt(2.0 / (embed_dim * kernel_size)) scale2 = jnp.sqrt(2.0 / (hidden_dim * kernel_size)) params = { 'conv1_w': jr.normal(keys[0], (kernel_size, embed_dim, hidden_dim)) * scale1, 'conv1_b': jnp.zeros(hidden_dim), 'conv2_w': jr.normal(keys[1], (kernel_size, hidden_dim, hidden_dim)) * scale2, 'conv2_b': jnp.zeros(hidden_dim), 'linear_w': jr.normal(keys[2], (hidden_dim, 1)) * jnp.sqrt(2.0 / hidden_dim), 'linear_b': jnp.zeros(1), } return params def duration_predictor(params, x): """从音素嵌入预测对数时长。x: (batch, seq, embed)。""" # 卷积层 1 配 ReLU h = jax.lax.conv_general_dilated( x.transpose(0, 2, 1), # (batch, embed, seq) params['conv1_w'].transpose(2, 1, 0), # (out, in, kernel) window_strides=(1,), padding='SAME' ).transpose(0, 2, 1) + params['conv1_b'] # 回到 (batch, seq, hidden) h = jax.nn.relu(h) # 卷积层 2 配 ReLU h = jax.lax.conv_general_dilated( h.transpose(0, 2, 1), params['conv2_w'].transpose(2, 1, 0), window_strides=(1,), padding='SAME' ).transpose(0, 2, 1) + params['conv2_b'] h = jax.nn.relu(h) # 线性投影到标量 log_dur = (h @ params['linear_w'] + params['linear_b']).squeeze(-1) return log_dur # 损失:对数时长上的 MSE(FastSpeech 中的标准做法) def loss_fn(params, embeddings, durations, masks): log_dur_pred = duration_predictor(params, embeddings) log_dur_true = jnp.log(jnp.clip(durations, 1.0, None)) sq_err = (log_dur_pred - log_dur_true) ** 2 * masks return jnp.sum(sq_err) / jnp.sum(masks) grad_fn = jax.jit(jax.value_and_grad(loss_fn)) # 训练循环 params = init_duration_predictor(jr.PRNGKey(0)) lr = 1e-3 losses = [] for epoch in range(300): loss_val, grads = grad_fn(params, embeddings, durations, masks) params = jax.tree.map(lambda p, g: p - lr * g, params, grads) losses.append(float(loss_val)) # 在一个样本上评估 log_dur_pred = duration_predictor(params, embeddings[:1]) dur_pred = jnp.exp(log_dur_pred[0]) dur_true = durations[0] mask = masks[0] valid_len = int(jnp.sum(mask)) fig, axes = plt.subplots(1, 2, figsize=(14, 5)) axes[0].plot(losses, color='#3498db', linewidth=1.5) axes[0].set_xlabel('Epoch') axes[0].set_ylabel('MSE Loss (log-duration)') axes[0].set_title('Duration Predictor Training') axes[0].set_yscale('log') x_pos = jnp.arange(valid_len) width = 0.35 axes[1].bar(x_pos - width/2, dur_true[:valid_len], width, color='#27ae60', label='Ground truth', alpha=0.8) axes[1].bar(x_pos + width/2, dur_pred[:valid_len], width, color='#e74c3c', label='Predicted', alpha=0.8) axes[1].set_xlabel('Phoneme index') axes[1].set_ylabel('Duration (frames)') axes[1].set_title('Duration Prediction vs Ground Truth') axes[1].legend() plt.tight_layout() plt.show()
import jax import jax.numpy as jnp import jax.random as jr import matplotlib.pyplot as plt def init_residual_block(key, channels, kernel_size, dilation): """初始化一个膨胀残差卷积块。""" k1, k2 = jr.split(key) scale = jnp.sqrt(2.0 / (channels * kernel_size)) return { 'conv1_w': jr.normal(k1, (kernel_size, channels, channels)) * scale, 'conv1_b': jnp.zeros(channels), 'conv2_w': jr.normal(k2, (kernel_size, channels, channels)) * scale, 'conv2_b': jnp.zeros(channels), 'dilation': dilation } def residual_block(params, x): """x: (batch, time, channels)。带 LeakyReLU 的膨胀卷积残差块。""" h = jax.nn.leaky_relu(x, negative_slope=0.1) # 简化:使用标准卷积(膨胀在概念上处理) h = jax.lax.conv_general_dilated( h.transpose(0, 2, 1), params['conv1_w'].transpose(2, 1, 0), window_strides=(1,), padding='SAME', rhs_dilation=(params['dilation'],) ).transpose(0, 2, 1) + params['conv1_b'] h = jax.nn.leaky_relu(h, negative_slope=0.1) h = jax.lax.conv_general_dilated( h.transpose(0, 2, 1), params['conv2_w'].transpose(2, 1, 0), window_strides=(1,), padding='SAME' ).transpose(0, 2, 1) + params['conv2_b'] return x + h def init_generator(key, n_mels=80, upsample_rates=(8, 8, 4), channels=128): """初始化一个最小的 HiFi-GAN 风格生成器。""" keys = jr.split(key, 10) params = {} # 输入投影:梅尔 bin -> 通道 params['input_w'] = jr.normal(keys[0], (7, n_mels, channels)) * 0.02 params['input_b'] = jnp.zeros(channels) # 上采样块(转置卷积) in_ch = channels for i, rate in enumerate(upsample_rates): k_size = rate * 2 scale = jnp.sqrt(2.0 / (in_ch * k_size)) out_ch = in_ch // 2 params[f'up{i}_w'] = jr.normal(keys[i+1], (k_size, in_ch, out_ch)) * scale params[f'up{i}_b'] = jnp.zeros(out_ch) # 每个尺度的残差块 params[f'res{i}_0'] = init_residual_block(jr.fold_in(keys[i+4], 0), out_ch, 3, 1) params[f'res{i}_1'] = init_residual_block(jr.fold_in(keys[i+4], 1), out_ch, 3, 3) in_ch = out_ch # 输出投影到单声道波形 params['output_w'] = jr.normal(keys[8], (7, in_ch, 1)) * 0.02 params['output_b'] = jnp.zeros(1) params['upsample_rates'] = upsample_rates return params def generator_forward(params, mel): """mel: (batch, time, n_mels) -> waveform: (batch, time * prod(rates), 1)。""" # 输入投影 h = jax.lax.conv_general_dilated( mel.transpose(0, 2, 1), params['input_w'].transpose(2, 1, 0), window_strides=(1,), padding='SAME' ).transpose(0, 2, 1) + params['input_b'] for i, rate in enumerate(params['upsample_rates']): h = jax.nn.leaky_relu(h, negative_slope=0.1) # 通过转置卷积上采样 k_size = rate * 2 h = jax.lax.conv_transpose( h.transpose(0, 2, 1), params[f'up{i}_w'].transpose(2, 1, 0), strides=(rate,), padding='SAME' ).transpose(0, 2, 1) + params[f'up{i}_b'] # 残差块 h = residual_block(params[f'res{i}_0'], h) h = residual_block(params[f'res{i}_1'], h) h = jax.nn.leaky_relu(h, negative_slope=0.1) out = jax.lax.conv_general_dilated( h.transpose(0, 2, 1), params['output_w'].transpose(2, 1, 0), window_strides=(1,), padding='SAME' ).transpose(0, 2, 1) + params['output_b'] return jnp.tanh(out) # 创建合成梅尔频谱图(模拟元音) n_mels = 80 n_frames = 50 mel = jnp.zeros((1, n_frames, n_mels)) # 在低频梅尔 bin 中加能量(模拟共振峰) mel = mel.at[:, :, 5:15].set(1.0) mel = mel.at[:, :, 20:25].set(0.6) # 初始化并运行生成器 key = jr.PRNGKey(42) params = init_generator(key, n_mels=n_mels, upsample_rates=(8, 8, 4), channels=128) waveform = generator_forward(params, mel) print(f"Input mel shape: {mel.shape}") print(f"Output waveform shape: {waveform.shape}") print(f"Upsample factor: {8 * 8 * 4} = {8*8*4}x") fig, axes = plt.subplots(2, 1, figsize=(12, 6)) axes[0].imshow(mel[0].T, aspect='auto', origin='lower', cmap='magma') axes[0].set_title('Input Mel Spectrogram') axes[0].set_ylabel('Mel bin') axes[0].set_xlabel('Frame') waveform_np = waveform[0, :, 0] axes[1].plot(waveform_np[:2000], color='#9b59b6', linewidth=0.5) axes[1].set_title('Generator Output Waveform (untrained - random noise)') axes[1].set_ylabel('Amplitude') axes[1].set_xlabel('Sample') plt.tight_layout() plt.show() print("Note: The output is noise because the generator is untrained.") print("In practice, adversarial + mel loss training shapes this into speech.")
import jax import jax.numpy as jnp import jax.random as jr import matplotlib.pyplot as plt # 生成带有语音/静音标签的合成对数梅尔能量特征 def generate_vad_data(key, n_sequences=100, n_frames=200, n_features=40): """模拟对数梅尔特征:语音区域能量更高且有结构。""" keys = jr.split(key, 5) all_features = [] all_labels = [] for i in range(n_sequences): k = jr.fold_in(keys[0], i) k1, k2, k3 = jr.split(k, 3) # 随机的语音/静音模式 label = jnp.zeros(n_frames) n_segments = jr.randint(k1, (), 2, 6) for seg in range(int(n_segments)): start = jr.randint(jr.fold_in(k2, seg), (), 0, n_frames - 20) length = jr.randint(jr.fold_in(k3, seg), (), 10, 50) end = jnp.minimum(start + length, n_frames) label = label.at[int(start):int(end)].set(1.0) # 特征:语音帧能量更高 + 有频谱结构 noise = jr.normal(jr.fold_in(keys[1], i), (n_frames, n_features)) * 0.3 speech_pattern = jnp.outer(label, jnp.exp(-jnp.arange(n_features) / 15.0)) features = speech_pattern * 2.0 + noise + 0.1 all_features.append(features) all_labels.append(label) return jnp.stack(all_features), jnp.stack(all_labels) key = jr.PRNGKey(123) features, labels = generate_vad_data(key) train_features, train_labels = features[:80], labels[:80] test_features, test_labels = features[80:], labels[80:] # 简单的基于 GRU 的 VAD 模型 def init_vad_model(key, input_dim=40, hidden_dim=64): keys = jr.split(key, 6) scale_ih = jnp.sqrt(2.0 / input_dim) scale_hh = jnp.sqrt(2.0 / hidden_dim) return { 'W_z': jr.normal(keys[0], (input_dim, hidden_dim)) * scale_ih, 'U_z': jr.normal(keys[1], (hidden_dim, hidden_dim)) * scale_hh, 'b_z': jnp.zeros(hidden_dim), 'W_r': jr.normal(keys[2], (input_dim, hidden_dim)) * scale_ih, 'U_r': jr.normal(keys[3], (hidden_dim, hidden_dim)) * scale_hh, 'b_r': jnp.zeros(hidden_dim), 'W_h': jr.normal(keys[4], (input_dim, hidden_dim)) * scale_ih, 'U_h': jr.normal(keys[5], (hidden_dim, hidden_dim)) * scale_hh, 'b_h': jnp.zeros(hidden_dim), 'W_out': jr.normal(jr.fold_in(keys[0], 99), (hidden_dim, 1)) * 0.1, 'b_out': jnp.zeros(1), } def gru_step(params, h, x): """单个 GRU 步。""" z = jax.nn.sigmoid(x @ params['W_z'] + h @ params['U_z'] + params['b_z']) r = jax.nn.sigmoid(x @ params['W_r'] + h @ params['U_r'] + params['b_r']) h_tilde = jnp.tanh(x @ params['W_h'] + (r * h) @ params['U_h'] + params['b_h']) h_new = (1 - z) * h + z * h_tilde return h_new def vad_forward(params, x): """x: (batch, time, features) -> logits: (batch, time)。""" batch_size, n_frames, _ = x.shape hidden_dim = params['W_z'].shape[1] h = jnp.zeros((batch_size, hidden_dim)) outputs = [] for t in range(n_frames): h = gru_step(params, h, x[:, t, :]) logit = (h @ params['W_out'] + params['b_out']).squeeze(-1) outputs.append(logit) return jnp.stack(outputs, axis=1) def bce_loss(params, features, labels): """VAD 的二元交叉熵损失。""" logits = vad_forward(params, features) probs = jax.nn.sigmoid(logits) probs = jnp.clip(probs, 1e-7, 1 - 1e-7) loss = -(labels * jnp.log(probs) + (1 - labels) * jnp.log(1 - probs)) return jnp.mean(loss) grad_fn = jax.jit(jax.value_and_grad(bce_loss)) # 训练 params = init_vad_model(jr.PRNGKey(0)) lr = 5e-3 losses = [] for epoch in range(200): loss_val, grads = grad_fn(params, train_features, train_labels) params = jax.tree.map(lambda p, g: p - lr * g, params, grads) losses.append(float(loss_val)) if epoch % 50 == 0: print(f"Epoch {epoch}: loss = {loss_val:.4f}") # 在测试集上评估 test_logits = vad_forward(params, test_features) test_preds = (jax.nn.sigmoid(test_logits) > 0.5).astype(jnp.float32) accuracy = jnp.mean(test_preds == test_labels) print(f"\nTest accuracy: {accuracy:.4f}") # 可视化一个测试样本 idx = 0 fig, axes = plt.subplots(3, 1, figsize=(14, 7)) axes[0].imshow(test_features[idx].T, aspect='auto', origin='lower', cmap='magma') axes[0].set_title('Log-Mel Energy Features') axes[0].set_ylabel('Mel bin') axes[1].fill_between(range(200), test_labels[idx], alpha=0.4, color='#27ae60', label='Ground truth') axes[1].plot(jax.nn.sigmoid(test_logits[idx]), color='#e74c3c', linewidth=1.5, label='Predicted probability') axes[1].axhline(0.5, color='gray', linestyle='--', linewidth=0.8) axes[1].set_ylabel('Speech probability') axes[1].legend() axes[1].set_title('VAD Predictions') axes[2].fill_between(range(200), test_labels[idx], alpha=0.4, color='#27ae60', label='Ground truth') axes[2].fill_between(range(200), test_preds[idx], alpha=0.4, color='#f39c12', label='Predicted (threshold=0.5)') axes[2].set_ylabel('Speech / Silence') axes[2].set_xlabel('Frame') axes[2].legend() axes[2].set_title('VAD Binary Decision') plt.tight_layout() plt.show()