LSTM前向计算
基本算子
LSTM模型的运算公式和求导已经在Neural Networks CheatSheet一文中有所总结,从基本算子的视角来看,大致可以描述为下图所示:

考虑到计算并行性,可以直接把四个权重矩阵拼接成一个来作为大的权重矩阵,\(x_t\)与之相乘以后将结果Slice成四份,即为四个Cell的向量表示。同样地,对\(h_{t-1}\)也可以进行同样的权重合并,两者结果相加以后再Slice即可,这样就完成了一个基本Cell的主要运算过程,后续的操作其实就是对不同Cell向量进行相应的Element-wise操作了。
上述计算过程称为LSTM的一个Node,那么我们考虑这个Node的输入和输出,可以看到,它接受上一个Node传递过来的一对状态\([h_{t-1}, c_{t-1}]\),以及当前的输入向量\(x_t\),然后产生一对输出状态\([h_t, c_t]\)。而在输出状态中,\(h_t\)就是当前Node的实际输出(向量),也可以记为\(y_t\);而\(c_t\)输出给外部,因此它不影响代价函数的计算,只是作为隐藏状态传递给下一个Node。
MXNet实现
对于一个接受输入\(x_t\), \(h_{t-1}\), \(c_{t-1}\),经过内部一套运算后输出\(h_t\), \(c_t\)的过程,它的计算图实际上是很容易定义的,MXNet的实现如下:
:::Python
def __call__(self, inputs, states):
i2h = symbol.FullyConnected(data=inputs, weight=self._iW, bias=self._iB,
num_hidden=self._num_hidden*4,
name='%si2h'%name)
h2h = symbol.FullyConnected(data=states[0], weight=self._hW, bias=self._hB,
num_hidden=self._num_hidden*4,
name='%sh2h'%name)
gates = i2h + h2h
slice_gates = symbol.SliceChannel(gates, num_outputs=4,
name="%sslice"%name)
in_gate = symbol.Activation(slice_gates[0], act_type="sigmoid",
name='%si'%name)
forget_gate = symbol.Activation(slice_gates[1], act_type="sigmoid",
name='%sf'%name)
in_transform = symbol.Activation(slice_gates[2], act_type="tanh",
name='%sc'%name)
out_gate = symbol.Activation(slice_gates[3], act_type="sigmoid",
name='%so'%name)
next_c = symbol._internal._plus(forget_gate * states[1], in_gate * in_transform,
name='%sstate'%name)
next_h = symbol._internal._mul(out_gate, symbol.Activation(next_c, act_type="tanh"),
name='%sout'%name)
return next_h, [next_h, next_c]
可以看到,这基本上就是按照上述计算流程翻译为Python而已。需要注意的一点是,上述计算过程实际上只是定义在计算图上的抽象计算,它的效果是建立相应的图结构,而并不是真的执行计算。那么抽象计算与真实计算的分界线在哪里?什么样的计算不能表述为抽象计算图?我个人的理解是,只要计算不依赖于符号的实际取值,那么这个计算过程就能够被静态计算图所描述。在上图中我们看到,我们已知做Slice操作的时候是要把那个向量切成四份,因此这个操作不依赖于数据的实际值,可以直接这样用计算图中的节点来描述。
静态计算图实现LSTM
Loop Unroll
当实现了单个Node的计算以后,只要给定输入序列\([x_1, x_2, ..., x_t]\)的长度,我们就可以将这么多个Node节点首尾相连,从而用于计算这个输入所对应的输出序列\([y_1, y_2, ..., y_t]\)。LSTM的展开大体如下图所示:

这个过程看起来是十分简单,只需要将上一个节点的输出作为下一次的输入,外加相应的数据即可。并且,这是一个对RNN来说通用的抽象,所以在MXNet中被实现为基类方法:
:::Python
def unroll(self, length, inputs, begin_state=None, layout='NTC', merge_outputs=None):
self.reset()
inputs, _ = _normalize_sequence(length, inputs, layout, False)
if begin_state is None:
begin_state = self.begin_state()
states = begin_state
outputs = []
for i in range(length):
output, states = self(inputs[i], states)
outputs.append(output)
outputs, _ = _normalize_sequence(length, outputs, layout, merge_outputs)
return outputs, states
从上述代码中可以看到,首先我们需要得到初始的输入状态(通常就是全零),来输入给第一个Node,随后就可以不断把当前Node的输出作为下一个Node的输入了,这个循环展开过程是显而易见的。并且,每一次循环也要注意保存当前的输出,以用于最后整体计算代价函数,误差反传。这部分代码整体来说是很好理解的。
当我们这样以循环展开的方式来定义LSTM以后,它就由一个带有循环性质的动态结构变成了静态计算图结构。可想而知,此时我们也就不再需要BPTT算法来对它进行训练了。因为计算图是确定的,计算过程是有限的,所以只要按照自动求导的原理倒着这样过一遍即可,效果与BPTT一定是等价的。
Stack Multiple RNN Cells
由于RNN的输入是时间序列,输出也是时间序列,因此一种高级的玩法是将多层Cell叠加在一起,例如多层LSTM叠加,从而有效提高模型的表达能力。那么在静态计算图框架下,对于这种模型的静态图展开也非常容易,只需要遍历所有Cell,一个一个把它们展开,然后上一个Cell的输出作为下一个Cell的输入即可,本质上跟单个Cell的循环展开是一样的。
MXNet的代码实现如下:
:::Python
def unroll(self, length, inputs, begin_state=None, layout='NTC', merge_outputs=None):
self.reset()
num_cells = len(self._cells)
if begin_state is None:
begin_state = self.begin_state()
p = 0
next_states = []
for i, cell in enumerate(self._cells):
n = len(cell.state_info)
states = begin_state[p:p+n]
p += n
inputs, states = cell.unroll(length, inputs=inputs, begin_state=states, layout=layout,
merge_outputs=None if i < num_cells-1 else merge_outputs)
next_states.extend(states)
return inputs, next_states
可以看到这个代码结构跟上面的展开单个Cell本质上也没太大区别,只不过这里是要获取所有Cell的初始状态,并且一个重要区别在于只把上一层的输出连接到下一层的输入,而上一层最后输出的额外状态如LSTM的\([h_t, c_t]\)就直接扔掉不要了。补充一点,级联RNN之所以能work,是因为上一层网络先对输入序列做了特征提取和变换,然后把结果序列输入给下一层去处理,提取更加上层的特征。
Sequence Buckets
处理完循环展开以后,此时我们遇到另外一个问题,那就是LSTM处理的是可变长度的序列,但我们的静态计算图一旦定义,就只能处理固定长度的序列了。于是最暴力的解决方案就是,干脆为每一个序列长度都创建一个静态计算图,然后在训练过程中所有子图的权重都是共享的,误差反传的时候更新同一个权重矩阵即可。这样做当然可行,但是资源消耗未免过于夸张,一个更靠谱的做法是先对所有出现过的序列长度聚类一下,找出几个典型的长度,然后把那些不零不整的序列pad到这些长度上面就行了。下面引用Yangqing Jia大神对此的描述:
一个折衷的方案就是,对于sequence先做聚类,预设几个固定的长度bucket,然后每个sequence都放到它所属的bucket里面去,然后pad到固定的长度。这样一来,首先我们不需要折腾while loop了,每一个bucket都是一个固定的computation graph;其次,每一个sequence的pad都不是很多,对于计算资源的浪费很小;再次,这样的实现很简单,就是一个给长度聚类,对于framework的要求很低。
另外,既然有padding那么当然要对padding的元素进行一下特殊处理。在NLP里面常见的方案是为padding元素分配一个特殊的label例如-1,然后在计算loss的时候忽略掉这个label就行了。