Force alignment(强制对齐)是语音任务中的一个常用模块,在帧级别的语音识别和语音合成任务中应用广泛,同时也是字幕自动打轴、口语评测等任务中的核心算法。强制对齐和识别解码有一定的相似性,识别解码是指给定一个语音识别模型,输入一个语音序列,输出对应的识别文本。而强制对齐是指给定一个语音识别模型,输入一个语音序列和一个正确的标注文本(可以是文字序列,也可以是音素序列等),给出语音帧和文本的对应关系,如第20帧到第25帧对应着'a'这个音素。

近日,WeNet 增加了 CTC alignment 功能。由于 WeNet 使用了如下的 Joint CTC/AED 结构,可以很方便的利用 CTC decoder 部分来完成对齐功能。即对于一个训练好的模型,给定的一段音频和对应的标注,通过 CTC decoder 部分的得分信息来获得语音和文本的对齐结果。对齐的单元跟建模单元有关,由于WeNet 中文模型使用了字级别的端到端的方法,因此得到的是字级别的对齐。如果需要 phone(音素)级别的对齐功能,需要训练一个 phone 级别的模型。



对齐

我们首先从建模的角度理解一下对齐。语音识别任务,需要对输入音频序列 X = [x1,x2,x3...,xt...,xT] (通常是 fbank 或 mfcc 等音频特征)和输出的标注数据文本序列 Y = [y1,y2,y3...,yu...,yU] 关系进行建模,其中 X 的长度一般大于 Y 的长度。如果能够知道yu和xt的对应关系,就可以将这类任务变成语音帧级别上的分类任务,即对每个时刻 xt 进行分类得到 yu。为了得到输入输出的对应关系,通常需要对输出序列进行扩展,扩展为和输入序列一样长。而不同的建模方法采取的扩展方法也不一样,在语音识别任务这种输入序列长度大于输出序列长度的任务中,通常包括对输出序列的单元进行复制和插入占位符 blank 两种操作。HMM 使用了复制的方式,CTC 则二者同时使用,RNN-T 则只使用了插入占位符 blank 的方式。下图来源自李宏毅老师的课程的 ppt,可以生动说明这三者做语音识别任务对齐的区别。


CTC 模型用下面的算法产生所有可能的扩展序列,其中每一个 CTC 扩展序列就是一个 CTC 对齐。其中 N 为输出序列长度,T 为输入的序列长度。(注为了简洁,以下未考虑特殊情况)

产生 b0 个blank 占位符
For n = 1 to N
    产生 第 n 个token tn 次
    产生 blank 占位符  bn 次
其中 tn 与 bn 需要满足如下限制:
b0 + t1 + b1 + t2 + b2 + ... tn + bn = T


再次通过下图来更好的理解产生所有可能的扩展序列的过程。


对于训练任务来说,在使用 CTC 目标训练时,本质上就是考虑上述产生所有可能的 CTC 对齐,并把每一种对齐情况下分类损失加起来作为目标函数。在实际实现时,并不会真的穷举所有对齐,而是利用一个高效算法进行计算。

而对齐任务来说,我们目的是寻找一个概率最大 CTC 扩展序列,最简单的方法是使用训练好的模型对所有可能的对齐路径进行打分,选择一条最高得分的路径即可,但是这样穷举的做法时间复杂度是指数级别的。这时可以通过维特比算法,来降低时间复杂度,解决这个问题。


维特比算法

Viterbi(维特比)算法是个动态规划的算法,Viterbi 算法可以得到一条概率最大的回溯路径,而回溯路径就是我们需要的对齐。对于维特比算法可以参考https://www.zhihu.com/question/20136144进行理解和学习。维特比算法的基础可以概括为下面三点(来源于吴军:数学之美):

  1. 如果概率最大的路径经过篱笆网络的某点,则从起始点到该点的子路径也一定是从开始到该点路径中概率最大的。
  2. 假定第 t 时刻有 k 个状态,从开始到 t 时刻的 k 个状态有 k 条最短路径,而最终的最短路径必然经过其中的一条。
  3. 根据上述性质,在计算第 t+1 时刻的最短路径时,只需要考虑从开始到当前的k个状态值的最短路径和当前状态值到第 t+1 时刻的最短路径即可。如求t=3时的最短路径,等于求 t=2 时,从起点到当前时刻的所有状态结点的最短路径加上 t=2 到 t=3 的各节点的最短路径。

基于此,通常 HMM 维特比算法包括以下流程:


简单的说维特比算法其实是求解多步骤,并且每步都进行多选择模型的这一类最优选择问题。对于每一步的所有可能的选择,维特比算法都保存了他们前续所有步骤到当前步骤当前选择的最小总代价(或者最大价值)以及当前代价的情况下前一步骤的选择。依次计算完所有步骤后,通过回溯的方法不断找寻前一步骤的选择即可找到完整的最优选择路径。


CTC 维特比算法

CTC 与 HMM 等一样,均可以使用维特比算法来求解最优的路径。对于 CTC 语音识别任务来说;多步骤相当于时间 t,每步骤的选择相当于每个时刻可选择的状态,如果是字建模的话,则表示每个时刻可选择的字或占位符 blank。WeNet 中的 CTC alignment 在使用维特比算法进行求解时也使用了上述的方法。但是由于 CTC 引入了占位符 blank,因此在使用 CTC 在维特比算法进行求解时状态跳转的处理与经典的 HMM 进行维特比算法时各流程的细节不完全一致。

  • 初始化对于初始化来说,由于 CTC 引入了 blank ,因此其第一时刻的初始状态可能为 blank 也可能为标注的 token 序列的第一个 token 如下图左上角的两个蓝色点。

  • 递推对于递推过程来说,CTC 由于引入了 blank,导致其状态跳转与 HMM 不相同。对于经典 HMM 来说,当前时刻 t 的状态s,可由 t-1 时刻任意状态跳转得到。而对于 CTC 来说则情况略微复杂些。根据状态是否为 blank主要可分为两类情况:

    • a. t 时刻的状态 s 为 blank,则状态 s 可由 t-1 时刻的 s 以及 t-1 时刻的 s-1 跳转而来
    • b. t 时刻的状态 s 为非 blank,这时有两种可能:

1)状态 s 与状态 s-2 相同,则状态 s 可由 t-1 时刻的 s 以及 t-1 时刻的 s-1 跳转而来,如 s 为下图的 x2 时刻的第二个 e 状态。这种情况的跳转与情况 a 相同,为此可以与情况 a 进行合并。


2)状态 s 与状态 s-2 不相同,则状态 s 可由 t-1 时刻的 s 、 t-1 时刻的 s-1 以及 t-1 时刻的 s-2 状态跳转而来,如下图的绿点表示的状态可由三个红色点表示的状态跳转而来。


  • 终止对于终止来说,CTC 可以由 blank,或实际标注的最后一个 token 作为终止状态(见上图的右下角两个蓝色点),而 HMM 只能实际标注的最后一个token 作为终止状态。

  • 最优路径回溯回溯的过程与 HMM 一致。


WeNet CTC alignment 的实现

接下来我们通过 WeNet 上的代码,再次理解一下 CTC 的维特比算法。

  • 数据处理

    将标注序列 y 插入 blank,如标注 y 为 c a t,插入 blank 占位符 ϵ 后为 ϵ c ϵ a ϵ t ϵ。ctc_probs 表示 CTC decoder 产生的概率分布。log_alpha[t, s] 记录了 t 时刻跳转至状态 s 的所有可能的路径中的最高得分。state_path[t, s] 则记录了 t 时刻的 s 状态由前一时刻的哪个最可能的状态跳转而来,用于最后的回溯。

y_insert_blank = insert_blank(y, blank_id)

log_alpha = torch.zeros((ctc_probs.size(0), len(y_insert_blank)))
log_alpha = log_alpha - float('inf')
state_path = (torch.zeros(
    (ctc_probs.size(0), len(y_insert_blank)), dtype=torch.int16) - 1
)


  • 初始化初始化 t0 时刻的 blank 以及 实际标注序列的第一个字符的概率(由CTC decoder 中的 softmax 输出的分布 ctc_probs 得到)
log_alpha[0, 0] = ctc_probs[0][y_insert_blank[0]]
log_alpha[0, 1] = ctc_probs[0][y_insert_blank[1]]


  • 递推递推过程首先根据 CTC 的跳转规则,对当前t时刻的任意状态s,找到其由上一时刻的哪些状态们(两个状态:分支1,或三个状态:分支2)跳转而来,这些状态由 prev_state 记录下来,对应的跳转至这些状态的路径们的最高得分由 candidates 记录。然后计算这些可能的状态跳转至当前状态s路径们(两条路径:分支1,或三条路径:分支2)的得分,选取得分最高得分作为的作为当前 t 时刻 s 状态的最高得分即 log_alpha[t, s],被选取的前一时刻的状态通过 state_path[t, s] 记录下来
for t in range(1, ctc_probs.size(0)):
    for s in range(len(y_insert_blank)):
        # 只能由上一时刻的相同状态或前一状态跳转而来的情况
        if y_insert_blank[s] == blank_id or s < 2 or y_insert_blank[
                s] == y_insert_blank[s - 2]:
            # 得到可跳转至前一时刻状态们的候选的路径的得分
            candidates = torch.tensor(
                [log_alpha[t - 1, s], log_alpha[t - 1, s - 1]])
            # 得到候选的前一时刻的状态
            prev_state = [s, s - 1]
        # 可由上一时刻的相同状态、前一状态、前前状态跳转而来的情况
        else:
            # 得到可跳转至前一时刻状态们的候选的路径的得分
            candidates = torch.tensor([
                log_alpha[t - 1, s],
                log_alpha[t - 1, s - 1],
                log_alpha[t - 1, s - 2],
            ])
            # 得到候选的前一时刻的状态
            prev_state = [s, s - 1, s - 2]
        # 记录可能跳转至s的路径们的最高得分
        log_alpha[t, s] = torch.max(candidates) + ctc_probs[t][y_insert_blank[s]]
        # 剪枝 只选取对于当前状态来说,最可能跳转过来的前一时刻的状态, 用于回溯
        state_path[t, s] = prev_state[torch.argmax(candidates)]
  • 终止最终的终止字符可能为 ϵ 或实际标注的最后一字符,比较最终时刻可跳转至最后一字符的所有路径的最高分 log_alpha[-1, len(y_insert_blank) - 2] 与 最终时刻跳转至实际标注的最后一字符的所有路径的最高分 log_alpha[-1, len(y_insert_blank) - 2],得分高的即为最终时刻的状态。
state_seq = -1 * torch.ones((ctc_probs.size(0), 1), dtype=torch.int16)

candidates = torch.tensor([
    log_alpha[-1, len(y_insert_blank) - 1],
    log_alpha[-1, len(y_insert_blank) - 2]
])
prev_state = [len(y_insert_blank) - 1, len(y_insert_blank) - 2]
state_seq[-1] = prev_state[torch.argmax(candidates)]


  • 回溯通过上一步确定的最后时刻的状态不断的回溯找到前一时刻的状态,直至找到第一时刻结束。
for t in range(ctc_probs.size(0) - 2, -1, -1):
    state_seq[t] = state_path[t + 1, state_seq[t + 1, 0]]


  • 最终输出
output_alignment = []
for t in range(0, ctc_probs.size(0)):
    output_alignment.append(y_insert_blank[state_seq[t, 0]])


目前 alignment 的入口可以通过 tools 目录下的 alignment.sh 脚本找到,或者通过 aishell/s0 目录下的 run.sh 的stage 7 找到。值得一提的是 WeNet 的 encoder 部分使用了卷积对输入进行了 subsample,因此在做维特比算法时是使用了 subsample 之后的时间序列。因此在使用 WeNet 产生的对齐时,可以根据自己的需求,通过简单的处理将对齐还原至 subsample 之前的对齐。此外,为了复用 WeNet 之前的对数据的读写,在做 alignment 之前也需将数据组织成与训练一致的格式。