WeNet更新支持了时间戳。解码器不仅可以返回 Nbest 解码结果,而且还可以返回其中每个字对应的时间信息。

在语音识别一些任务中,字级别的的时间戳和N-best 扮演着重要的作用。例如在视频应用中,语音识别结合字级别的时间戳可以在精确的时间显示字幕,在会议场景中,字级别的时间戳可以标定与会者在说某句话某个字的精确时间。N-best 则包含了更多的识别信息,并提升识别的下游任务,如纠错、NLP 等的准确性。

下面将详细介绍时间戳在WeNet 中的实现(时间戳和 N-best 仅在 runtime 中实现)。


CTC Prefix Beam Search

在介绍时间戳的实现之前,我们先来回顾一下CTC Prefix Beam Search 算法。


上图出自 Sequence ModelingWith CTC,相信大家都已经耳熟能详。神经网络输出一个T  X M 的矩阵,其中T 表示音频的帧数 (10 帧); M 表示词典的大小 (5 个字母)。CTC Prefix Beam Search 算法则在该矩阵的基础上,找出概率最高的N 条路径。假设模型的输出如下图左上角的表格所示:


CTC Prefix Beam Search的过程,每个时刻有如下3个动作:

1. 扩展:根据前缀串和当前时刻的输出,计算新串的概率。

2. 规约:将规约串相同的候选概率相加。

3. 裁剪:仅保留top k个最好的序列做下一时刻的拓展,绿色的表示保留,红色的表示被裁减掉,图中k为3。



时间戳

每个前缀串可以由多个串规约而成。WeNet使用前缀串的被规约串中最优的一条路径,即viterbi路径来记录时间信息,viterbi路径中记录了每个字峰值的时间。如下图所示:


解码后一共得到三个解码结果:a, ab和ba。

1. 对于解码结果a来说,考虑到剪枝策略,因此规约前的串只可能是εaε、εaa或者aaa。

· viterbi分数较高的是aaa,为 0.4 X 0.35 X 0.50 = 0.07。

· a的峰值在T = 3,概率为0.50,即时间戳为T = [3]。

2. 对于解码结果ab来说,规约前的串可能是aab或者aεb。

3. viterbi分数较高的是aεb,为0.40 X .040 X .040 = 0.64 。a和b的时间戳为T = [1, 3]。

4. 对于解码结果ba来说,规约前的串可能是εba、bεa、baε、baa或者bba。

· viterbi分数较高的是bεa,为0.35 X 0.40 X 0.50 = 0.07 。b和a的时间戳为T = [1, 3]。

通常一个字的时间戳信息应该包括起始时间和终止时间,而使用上述算法,我们只能获取该字峰值所在的时间。因此在WeNet的实现中,考虑到延迟等因素,我们将峰值所在的时间当做该字的终止时间,上一个字峰值所在的时间当做起始时间。


示例

在启动客户端的时候,可以通过nbest参数来让服务器返回多个候选结果和对应的时间戳。程序的运行结果如下图所示:

 


代码实现

WeNet使用HashMap来保存解码过程中产生的前缀串及其对应的分数信息。分数信息的结构体定义如下所示:

struct PrefixScore {
  float cur_token_prob = -kFloatMax;  // 当前 token 峰值的概率
  float s = -kFloatMax;               // 以 ε 结尾的分数
  float ns = -kFloatMax;              // 以非 ε 结尾的分数
  float v_s = -kFloatMax;             // 以 ε 结尾的 viterbi 分数
  float v_ns = -kFloatMax;            // 以非 ε 结尾的 viterbi 分数
  std::vector<int> times_s;           // 以 ε 结尾的 viterbi 路径的时间戳
  std::vector<int> times_ns;          // 以非 ε 结尾的 viterbi 路径的时间戳

  // 前缀串的分数为 s 和 ns 的和
  float score() const { return LogAdd(s, ns); }
  // viterbi 分数为 max(v_s, v_ns)
  float viterbi_score() const { return v_s > v_ns ? v_s : v_ns; }
  // 根据 viterbi 分数选择前缀串的时间戳
  const std::vector<int>& times() const {
    return v_s > v_ns ? times_s : times_ns;
  }
};


主要代码的实现在 decoder/ctc_prefix_beam_search.cc 的Search函数中。代码通过for循环遍历每一个时刻,获取每一个时刻的输出,然后执行CTC Prefix Beam Search的过程。代码主要分为四部分:

1. 第一次剪枝

2. Token Passing

3. 第二次剪枝

4. 更新前缀串


第一次剪枝

在上面表格中,词典只包含3个字母['ε', 'a', 'b'],因此每一时刻的输出都包含3个字母。而我们的字典一共包含4233个汉字,需要通过剪枝来降低计算的开销。这里 opts_.first_beam_size 默认的取值为10,即只保留概率最高的前10个汉字的概率及其索引。

// 1. First beam prune, only select topk candidates
std::tuple<Tensor, Tensor> topk = logp_t.topk(opts_.first_beam_size);
Tensor topk_score = std::get<0>(topk);
Tensor topk_index = std::get<1>(topk);


Token Passing

Token Passing部分的代码首先通过for循环遍历当前时刻的10个输出,然后对前缀串进行扩展和规约(代码如下):

// 2. Token Passing
// next_hyps 记录扩展规约后的前缀串,即下一个时刻的前缀串,避免更新当前时刻产生的前缀串
std::unordered_map<std::vector<int>, PrefixScore, PrefixHash> next_hyps;
for (int i = 0; i < topk_index.size(0); ++i) {
  int id = topk_index[i].item<int>();
  auto prob = topk_score[i].item<float>();
  for (const auto& it : cur_hyps_) {
    const std::vector<int>& prefix = it.first;
    const PrefixScore& prefix_score = it.second;
    // 如果 prefix 不在 next_hyps 中, next_hyps[prefix] 则会插入默认的分数信息
    if (id == opts_.blank) {
      // Case 0: *a + ε => *a; *aε + ε => *a
      // 当前时刻输出 ε,表示新串与前缀串相同
      PrefixScore& next_score = next_hyps[prefix];
      // 由于当前时刻可能已经产生了相同的新串,所以需要进行规约
      // 即新串以 ε 结尾的分数 next_score.s 为两者之和:
      //  1. 新串以 ε 结尾的分数 next_score.s
      //  2. 前缀串的分数 prefix_score.score() 和当前输出的概率的对数 prob 和
      next_score.s = LogAdd(next_score.s, prefix_score.score() + prob);
      // 新串以 ε 结尾的 viterbi 分数 next_score.v_s 为:
      // 前缀串的 viterbi 分数 prefix_score.viterbi_score() 和当前输出的概率 prob 的对数和
      next_score.v_s = prefix_score.viterbi_score() + prob;
      // 新串以 ε 结尾的 viterbi 路径的时间戳 next_score.times_s 等于:
      // 前缀串的时间戳 prefix_score.times()
      next_score.times_s = prefix_score.times();
    } else if (!prefix.empty() && id == prefix.back()) {
      // 前缀串不为空,且当前时刻的输出与前缀串最后一个字相同
      // 假设当前时刻的输出为 a,则上一时刻的输出可能是 a 或者 ε
      // Case 1: *a + a => *a
      // 当前时刻输出 a,表示新串与前缀串相同
      PrefixScore& next_score1 = next_hyps[prefix];
      // 由于新串可能已经存在 next_hyps 中,所以需要进行规约
      // 即新串以非 ε 结尾的分数 next_score1.ns 为两者之和:
      //  1. 新串以非 ε 结尾的分数 next_score1.ns
      //  2. 前缀串以 ε 结尾的分数 prefix_score.ns 和当前输出的概率 prob 的对数和
      // 新串以非 ε 结尾的分数 next_score1.ns 为:
      // 前缀串以非 ε 结尾的分数 prefix_score.ns 和当前输出的概率 prob 的对数和
      next_score1.ns = LogAdd(next_score1.ns, prefix_score.ns + prob);
      // 判断是否需要更新新串以非 ε 结尾的 viterbi 分数 next_score1.v_ns
      if (next_score1.v_ns < prefix_score.v_ns + prob) {
        next_score1.v_ns = prefix_score.v_ns + prob;
        // 判断是否需要更新新串中最后一个字峰值的概率 next_score1.cur_token_prob
        if (next_score1.cur_token_prob < prob) {
          next_score1.cur_token_prob = prob;
          // 新串以非 ε 结尾的 viterbi 路径的时间戳 next_score1.times_ns 等于:
          // 前缀串以非 ε 结尾的 viterbi 路径的时间戳 prefix_score.times_ns
          next_score1.times_ns = prefix_score.times_ns;
          CHECK_GT(next_score1.times_ns.size(), 0);
          // 更新新串中最后一个字峰值的位置 next_score1.times_ns.back()
          next_score1.times_ns.back() = abs_time_step_;
        }
      }
      // Case 2: *aε + a => *aa
      // 将当前时刻的输出拼接到前缀串上,得到新串
      std::vector<int> new_prefix(prefix);
      new_prefix.emplace_back(id);
      PrefixScore& next_score2 = next_hyps[new_prefix];
      // 由于当前时刻可能已经产生了相同的新串,所以需要进行规约
      // 即新串以非 ε 结尾的分数 next_score2.ns 为两者之和:
      //  1. 新串以非 ε 结尾的分数 next_score2.ns
      //  2. 前缀串以 ε 结尾的分数 prefix_score.s 和当前输出的概率 prob 的对数和
      next_score2.ns = LogAdd(next_score2.ns, prefix_score.s + prob);
      // 判断是否需要更新新串以非 ε 结尾的 viterbi 分数 next_score2.v_ns
      if (next_score2.v_ns < prefix_score.v_s + prob) {
        // 新串以非 ε 结尾的 viterbi 路径的时间戳 next_score2.times_ns 等于:
        // 前缀串以 ε 结尾的 viterbi 路径的时间戳 prefix_score.times_s,拼接上当前时间步 abs_time_step_
        next_score2.v_ns = prefix_score.v_s + prob;
        next_score2.cur_token_prob = prob;
        next_score2.times_ns = prefix_score.times_s;
        next_score2.times_ns.emplace_back(abs_time_step_);
      }
    } else {
      // Case 3: *a + b => *ab, *aε + b => *ab
      // 当前时刻的输出与前缀串最后一个字不同,将当前的输出拼接到前缀串上得到新串
      std::vector<int> new_prefix(prefix);
      new_prefix.emplace_back(id);
      PrefixScore& next_score = next_hyps[new_prefix];
      // 由于当前时刻可能已经产生了相同的新串,所以需要进行规约
      // 即新串以非 ε 结尾的分数 next_score.ns 为两者之和:
      //  1. 新串以非 ε 结尾的分数 next_score.ns
      //  2. 前缀串的分数 prefix_score.score() 和当前输出的概率 prob 的对数和
      next_score.ns = LogAdd(next_score.ns, prefix_score.score() + prob);
      // 判断是否需要更新新串以非 ε 结尾的 viterbi 分数 next_score.v_ns
      if (next_score.v_ns < prefix_score.viterbi_score() + prob) {
        next_score.v_ns = prefix_score.viterbi_score() + prob;
        // 更新前缀串最后一个字峰值的概率
        next_score.cur_token_prob = prob;
        // 新串以非 ε 结尾的 viterbi 路径的时间戳 next_score.times_ns 等于:
        // 前缀串 viterbi 路径的时间戳 prefix_score.times(),拼接上当前时间步 abs_time_step_
        next_score.times_ns = prefix_score.times();
        next_score.times_ns.emplace_back(abs_time_step_);
      }
    }
  }
}


第二次剪枝

第二次剪枝只保留分数最高的前N条路径(N-Best),便于后续的重打分。这里 opts_.second_beam_size默认的取值为10。

// 3. Second beam prune, only keep top n best paths
std::vector<std::pair<std::vector<int>, PrefixScore>> arr(next_hyps.begin(),
                                                          next_hyps.end());
int second_beam_size =
  std::min(static_cast<int>(arr.size()), opts_.second_beam_size);
std::nth_element(arr.begin(), arr.begin() + second_beam_size, arr.end(),
                 PrefixScoreCompare);
arr.resize(second_beam_size);
std::sort(arr.begin(), arr.end(), PrefixScoreCompare);


更新前缀串

将next_hyps中的新串更新到前缀串集合cur_hyps中,并且获取当前每个解码结果的分数等信息。

// 4. Update cur_hyps_ with next_hyps and get new result
cur_hyps_.clear();
hypotheses_.clear();
likelihood_.clear();
viterbi_likelihood_.clear();
times_.clear();
for (auto& item : arr) {
  // 更新前缀串
  cur_hyps_[item.first] = item.second;
  // 更新解码结果
  hypotheses_.emplace_back(std::move(item.first));
  // 更新每个解码结果的分数
  likelihood_.emplace_back(item.second.score());
  // 更新每个解码结果的 viterbi 分数
  viterbi_likelihood_.emplace_back(item.second.viterbi_score());
  // 更新每个解码结果的时间戳信息
  times_.emplace_back(item.second.times());
}


总结

上述内容就是CTC Prefix Beam Search算法和时间戳在WeNet中的实现,虽然CTC Prefix Beam Search的整个过程较为简单,但是需要在其中保留更多的路径信息,以获取每条规约后路径的时间戳。对这部分内容感兴趣的同学,可以参考WeNet中提供的单元测试进行调试与学习。

[0]. WeNet. https://github.com/mobvoi/wenet

[1]. Sequence Modeling With CTC. https://distill.pub/2017/ctc

[2]. CTC prefix beam search. https://robin1001.github.io/2020/12/11/ctc-search

[3]. CTC的Decode算法-Prefix Beam Search. http://placebokkk.github.io/asr/2020/02/01/asr-ctc-decoder.html