Featured image of post MANNs:记忆增强神经网络

MANNs:记忆增强神经网络

介绍记忆增强神经网络的控制器、记忆矩阵、读写机制与基本计算流程。

定义

MANNs(Memory-Augmented Neural Networks,记忆增强神经网络)不是某一个固定模型,而是一类带有可读写记忆模块的神经网络架构。

普通 RNN 或 LSTM 主要把信息保存在隐藏状态中;MANNs 则进一步引入独立的记忆矩阵,使模型能够保存和检索更多信息。

MANNs 不能简单理解为 LSTM 的扩展或一种特殊门控单元。LSTM 可以作为 MANNs 的控制器;MANNs 的关键特征是控制器之外还存在可寻址的记忆模块。

常见架构

  1. NTM(Neural Turing Machine,神经图灵机)
  2. DNC(Differentiable Neural Computer,可微神经计算机)
  3. Memory Networks

基本结构

以 NTM 为例,系统通常由以下部分组成:

  1. 控制器(Controller):处理当前输入和历史信息,可以使用 RNN、LSTM 或前馈神经网络。
  2. 记忆矩阵(Memory):由多个可寻址的记忆槽组成。
  3. 读头(Read Head):根据读权重从记忆矩阵中读取信息。
  4. 写头(Write Head):根据写权重擦除旧信息并写入新信息。

MANNs 基本结构

MANNs 读写机制

若图表示当前时刻完成读取后的输出,读向量应记为 $r_t$,而不是 $r_{t-1}$。不过,控制器在计算当前隐藏状态时通常接收的是上一步的读取结果 $r_{t-1}$。

一次计算的基本流程

控制器更新状态

在时刻 $t$,控制器接收当前输入 $x_t$、上一步的读取结果 $r_{t-1}$ 和上一步的隐藏状态 $h_{t-1}$,并计算:

$$ h_t=\operatorname{Controller}(x_t,r_{t-1},h_{t-1}) $$

若控制器是前馈神经网络,则不一定存在 $h_{t-1}$;具体形式取决于模型设计。

生成读写参数

控制器的隐藏状态 $h_t$ 会通过线性层映射为读写头所需的接口参数,包括查询向量 $k_t$、读写权重 $w_t^r$ 与 $w_t^w$、擦除向量 $e_t$ 和添加向量 $a_t$。

例如:

$$ k_t=W_kh_t+b_k $$

在标准 NTM 中,$k_t$ 作为读取键(read key)通过基于内容的寻址生成读权重:先计算 $k_t$ 与各记忆槽 $M_t(i)$ 的余弦相似度,再经 Softmax 归一化:

$$ w_t^r(i)=\frac{\exp\left(\beta_t K(k_t,M_t(i))\right)}{\sum_j \exp\left(\beta_t K(k_t,M_t(j))\right)},\qquad K(u,v)=\frac{u\cdot v}{|u|\cdot|v|} $$

其中,$K(\cdot,\cdot)$ 是余弦相似度,$\beta_t$ 是寻址强度,通常也由控制器线性映射得到;与 $k_t$ 越相似的记忆槽会被分配到越大的读权重。

$$ e_t=\sigma(W_eh_t+b_e) $$

$$ a_t=\tanh(W_ah_t+b_a) $$

为了简化计算,本文省略上述内容寻址,直接对控制器输出做 Softmax 得到各记忆槽的权重:

$$ z_t^r=W_rh_t+b_r,\qquad w_t^r=\operatorname{Softmax}(z_t^r) $$

$$ z_t^w=W_wh_t+b_w,\qquad w_t^w=\operatorname{Softmax}(z_t^w) $$

读写权重满足:

$$ w_t^r(i)\ge 0,\qquad \sum_i w_t^r(i)=1 $$

$$ w_t^w(i)\ge 0,\qquad \sum_i w_t^w(i)=1 $$

在标准 NTM 中,得到最终读写权重通常还要经过插值、循环移位和锐化等步骤,并不只是上面的 Softmax 式子那么简单。

从记忆矩阵读取

读取并不是只选择一个记忆地址,而是对所有记忆槽进行加权求和:

$$ r_t=\sum_i w_t^r(i)M_t(i) $$

其中,$i$ 是记忆槽的位置编号,$M_t(i)$ 是时刻 $t$ 的第 $i$ 个记忆槽,$w_t^r(i)$ 是该槽的读权重,$r_t$ 是当前读取结果。

例如,设四个记忆槽为:

$$ m_1=[1,0],\quad m_2=[3,1],\quad m_3=[8,5],\quad m_4=[2,7] $$

读权重为:

$$ w_t^r=[0.05,0.10,0.80,0.05] $$

则:

$$ \begin{aligned} r_t &=0.05m_1+0.10m_2+0.80m_3+0.05m_4\ &=[6.85,4.45] \end{aligned} $$

擦除与写入

写入新信息前,先使用擦除向量删除部分旧信息:

$$ \widetilde{M}_t(i)=M_{t-1}(i)\odot\left(\mathbf{1}-w_t^w(i)e_t\right) $$

其中,$\odot$ 表示逐元素乘法,$\mathbf{1}$ 是与 $e_t$ 维度相同的全 1 向量。

随后将添加向量写入记忆矩阵:

$$ M_t(i)=\widetilde{M}_t(i)+w_t^w(i)a_t $$

一次写操作可以概括为“先擦除,再添加”。

产生模型输出

模型可以结合控制器状态 $h_t$ 和当前读取结果 $r_t$ 产生最终输出:

$$ y_t=f\left(W_y[h_t;r_t]+b_y\right) $$

其中,$[h_t;r_t]$ 表示向量拼接,$f$ 是由任务决定的输出函数,例如 Softmax、Sigmoid 或恒等映射。

整体流程总结

  1. 输入 $x_t$ 与上一步读取结果 $r_{t-1}$ 进入控制器。
  2. 控制器计算隐藏状态 $h_t$。
  3. 根据 $h_t$ 生成寻址与读写所需的接口参数。
  4. 读头计算读权重并得到读取结果 $r_t$。
  5. 写头按照写权重擦除旧信息并添加新信息。
  6. 模型结合 $h_t$ 与 $r_t$ 生成输出 $y_t$。

不同 MANN 架构的读写顺序和寻址机制可能不同,上述流程是便于理解的通用化描述。