本文以 CTC Loss 计算为例,分析 k2.compose 和 k2.intersect 的区别和联系
如果已经很熟悉 k2.compose = k2.invert + k2.intersect, 建议直接滑到文末抽奖环节(大周末的,多休息不好么)
本文对应的colab, 对比了纯手算,pytorch, k2 三种 ctc_loss 计算方式:https://colab.research.google.com/drive/1pS1XK30BFhFJPAAY3ZmJ96qYIzwnYWuZ?usp=sharing
注意:
本文不明确区分 fst/fsa, 统一用 fst 表示,毕竟 fsa 也可以视作“输入输出 lablel” 一样的特殊 fst。 k2 中有多个 intersect 相关的函数,比如 intersect_dense, intersect_dense_pruned。如非明确指出,本文中说某一概念适用于 "k2.intersect" 函数是指这一概念适用于所有的 "k2.intersect_*" 函数 k2 中关于 fst 的操作大多数是一元函数(只处理一个 fst),比如 connect, invert, reverse 等等。compose / intersect 是为数不多的二元函数(处理两个fst)。
众所周知, 基于 k2 的 CTC Loss 计算流程如下图所示(还没有用过的可以跑一下文章开头的colab, 再回到本文)。在不同的环节用了两个不同的二元操作 k2.compose 和 k2.intersect_dense。本文将以此为例,分析何时用 compose,何时用 intersect。

1. fst 的图表示了什么信息
让我们回到刚过去没多久的高一,函数是一种特殊的映射,表述输入集合与输出集合之间的对应关系。(小编实在找不到当初的课本,脑子里也记不起来标准的定义,就大概是这么个意思吧。另外下文也不区分“映射”和“函数”,统一称为“函数”。)

fst 可以看作一种特殊的函数,它建立了输入集合和输出集合的对应关系。如下图所示,这个 fst 函数表示的是 CTC 建模中用到的规则:“多个连续的token只保留一个”。假设标注文本对应的序列为[ 1])

注意我们关注的对象是能够从起始状态一直到终止状态的路径,由于图中自环的存在,所以这样的对象必然是无限多个。我们列举部分如下(忽略到 label==-1 对应的边):
路径 输入序列 输出序列 去除blk 0
0->1->3 1 1 1
0->0->1->3 01 01 1
0->0->0->1->3 001 001 1
0->1->1->3 11 10 1
0->1->1->2->3 110 100 1至此,我们至少可以从以下三个角度来观察 fst:
各路径对应的输入序列构成的集合 各路径对应的输出序列构成的集合 输入序列/输出序列对应的函数关系
2. 什么是 compose
这个词可以理解为“复合”函数。比如定义如下:
y = f(x)
z = g(y)
z = u(x) = g(f(x))则函数u称为函数g和f的复合函数。注意此时函数f的输出对应的是函数g的输入。
如上文所述,我们可以把 fst 看成一个函数, 那么两个 fst 做 compose 的含义不难想象。我们从函数的角度类比一下:
fst_u = k2.compose(fst_f, fst_g)它大概的意思,应该是fst_f的输出对应 fst_g 的输入。即:
fst_u = fst_g(fst_f(input))3. 什么时候用 compose
如上文所述,当我们关注的是fst_f的输出作为fst_g的输入的时候,应该用compose。分析一下 k2.compose(ctc_graph, linear_fst) 是否符合这种关系。
ctc_graph 图如下所示,输入是连续重复字符合并前,输出是连续重复字符合并后(合并是通过保留第一个,其余全置0的方式实现, 即 state 1和state 2上面的自环)。

而 linear_fst 表示标注文本,如下图所示,自然是连续重复字符合并后。

所以ctc_graph输出刚好接 linear_fst 的输入。这两个 fst 用 k2.compose 正合适。
结合复合函数的性质,z = g(f(x))产生的最终函数关系是f的输入与g的输出之间的关系。或者说f的输出和g的输入被消掉了。所以
decoding_graph = k2.compose(ctc_graph, linear_fst)的输入是 ctc_graph 的输入,即连续重复字符合并前。而输出是 linear_fst 的输出, 即连续重复字符合并后。其结构如图所示,即本文第一小节分析过的图:

4. 什么是 intersect, 什么时候该用 intersect?
简而言之,intersect 是集合求交集的意思。
接下来分析为什么 decoding_graph 和 dense_fsa_vec 要做intersect而不是compose。正如前文所说,一个 fst 可以从它包含的输入序列和输出序列的角度去分析。
如下图所示:dense_fsa_vec 表示的是每个时刻各个 token 之间相互跳转的概率,是连续重复字符合并前, 所以和 decoding_graph 的输入保持一致。

此时decoding_graph 和 dense_fsa_vec 各自的特点如下:
decoding_graph:任意的输入序列,通过 decoding_graph 这个函数,对应的输出序列都刚好对应标注,即连续重复字符合并后。反过来讲,能够生成该标注的所有输入序列都已经囊括到 decoding_graph的输入序列集合里。做到了不重不漏。但是它没有各序列的概率。 dense_fsa_vec:给定任意输入序列,都可以算出来对应的概率,但是它不知道哪些输入序列才能够生成标注。
所以,各位看官老爷们,现在decoding_graph 和 dense_fsa_vec 该如何协作简直是拍到了小编的脸上。一个知道怎么走,不知道概率;一个知道概率,但是不知道怎么走!直接对二者的输入序列做交集,然后互通有无,得到的结果就是既知道怎么走又知道对应的概率。
k2.intersect 就是干了这个事,所以二者 intersect 后,得到的结果如下图所示:

5. k2.compose = k2.invert + k2.intersect
结合上文,再回到 compose, 既然说 intersect 是对两个 fst 的输入序列做交集。那么 k2.compose(fst_f, fst_g) 是不是可以看成对 fst_f 的输出序列与 fst_g 的输入序列做交集。
compose 和 intersect 的核心反正都是求交集,不如共享一下代码。此时只要把 fst_f 的输入输出交换一下,得到inv_fst_f = fst_f.invert(), 那么待处理的问题就由:
fst_f 的**输出序列** 与 fst_g 的**输入序列**做交集转变为
inv_fst_f 的**输入序列** 与 fst_g 的**输入序列**做交集嗯,转变后就可以用 k2.intersect(invert_fst_f, fst_g)。等 intersect 完毕,记得再把输入输出的符号倒腾回来就行了。
k2.compose 的真实代码[1]关键步骤如下所示:
def compose(a_fsa: Fsa,
b_fsa: Fsa,
treat_epsilons_specially: bool = True,
inner_labels: Optional[str] = None) -> 'Fsa':
# 先取个 invert, 好复用 k2.intersect 函数
a_fsa_inv = a_fsa.invert()
# 此时的 aux_labels 其实是原来的"输入",保留住,等intersect 完还要倒腾回去。
a_fsa_inv.rename_tensor_attribute_('aux_labels', 'left_labels')
# 复用 intersect 函数
ans = intersect(a_fsa_inv, b_fsa, treat_epsilons_specially=treat_epsilons_specially)
# 把输入符号倒腾回去
ans.rename_tensor_attribute_('left_labels', 'labels')不知为何,小编此时想起来把大象关到冰箱里的三个步骤:
1. 冰箱门打开
2. 把大象赶紧去
3. 冰箱门关上k2.compose(fst_f, fst_g) 的实现也可以分成三步:
1. fst_f 取 invert 得到 inv_fst_f, 从而 invert_fst_f 的输入和 fst_g 的输入相对应。
2. k2.intersect(invert_fst_f, fst_g)
3. 把输入符号倒腾回去所以在 k2 中,最底层的 CUDA 的代码并没有 compose 的实现,只有 intersect 对应的代码实现。甚至底层的 intersect 的代码实现的时候,压根就没传各边对应的输出label,只传了两个 fst 各边对应的输入label。
6. 总结
本文简单分析了 k2.compose 与 k2.intersect 的区别与联系。
当关注的是输出与输入的关系时,应该用 k2.compose。此时更多的地把 fst 看作函数。 当关注的是输入与输入的关系时,应该用 k2.intersect。此时更多地把 fst 看作输入序列的集合。 k2.compose = k2.invert + k2.intersect
总之,横看成岭侧成峰,函数集合要分清。
参考资料
k2.compose 代码:https://github.com/k2-fsa/k2/blob/0a6cf0b7b5e4f58a3444cac85800f397f7f317f0/k2/python/k2/fsa_algo.py#L416
