笔记:Multi-Level Matching and Aggregation Network for Few-Shot Relation Classification


Multi-Level Matching and Aggregation Network for Few-Shot Relation Classification

作者:Ye et al., ACL 2019

目录

  • 简介
  • 方法
  • 实验结果
  • 总结

1 简介

本文还是在FewRel做小样本关系抽取,主要在ProtoCNN的基础之上, 在query instance 、 support set instances的embedding以及prototype 的计算上做了改进,引入交互式的方法--即考虑query和support set instances之间local、instance level的信息匹配,来encode query instance和support set instances,不同于之前的prototypical networks只是分别独立计算query和support set instance以及prototype。

2 方法

他算法图画得过于抽象,我就不放了。MLMAN 是在 ProtoCNN 的基础上修改的,即当 CNN 将句子和 position 编码后,得到 support instance 的表示\(C \in R^{T_s \times d_c}\) 和 query instance 的表示 \(Q \in R^{T_q \times d_c}\),之后,引入 MLMA。

(式子里 \(d_c\) 是 embedding size,\(T_q\) 是 query 句子长度,\(T_s=\Sigma_{k=1}^KT_k\) 是所有 support 句子横向拼接的长度。)\(^{[2]}\)

引入的MLMA主要分为三个部分Local Matching and Aggregation、Instance Matching and Aggregation 和Class Matching.

2.1 任务

类似传统prototypical networks,任务仍然是最小化目标函数,如下公式Eq (1) Eq (2).

其中\(S = \{s_k^i;i=1,...,N,k=1,...,K\}\)为train-support set且\(s_k^i\)为类别\(i\)的第\(k\)个instance。\(Q=\{(q_j,l_j);j=1,...,R\}\)为train-query set且\(R\)为样本的数目(在N类中剩余样本中随机采样),\(l_j\in\{1,...,N\}\)是instance \(q_j\)的label。

函数\(f(\{s_k^i\}_{k=1}^K,q)\)计算query instance 与每个类别的support instances \(\{s_k^i\}_{k=1}^K\)的匹配程度,我们主要关注的就是这个函数的设计。

2.2 Local Matching and Aggregation

这部分着实没太看懂不知道为什么这么设计,好像还涉及NLI的东西着实没看过,当然这个方法也是参考了一篇paper,感兴趣可以看看原文,这里直接引用另一篇优秀的博客\(^{[2]}\)

先计算 support 和 query 分别的 matched representations \(\tilde{Q}\)\(\tilde{C}\)~:

这里右下角标都指大矩阵中的某一列,长度都是 CNN embedding size。α 计算的是 token-level 相似度,\(\tilde{q}_m\) 就表示,作为 query 的第 m 个 token,和 support 交互的结果,这个结果衡量了 token m 与 seppot 整体的相似度(?),也就是 match 的含义所在。(其实我对这个式子的设计感觉非常迷惑,我编不出来有意义的解释 Orz)

得到 matched representations 后,就用各种暴力的方式,和原来的 embedding 拼接、求差、按位相乘,再 linear + relu + Bi-LSTM,再接一层特殊的 pooling,得到最后的 embedding。即先一顿操作:

然后用 \(\hat{Q}\) 过 Bi-LSTM 得到 \(\tilde{Q}\)(C 也是)。再过一个特殊的 pooling:

其中 \(\hat{S^k}\) 是把 \(\hat{C}\) 按照句子拆开得到的。

以上就是 local 的交互。local 在句子里就应该是表示 token 和周围邻居的交互,可是这里是 token 和对面所有句子的交互,这完全不 localized。(当然,他这个做法是有道理的,确实没必要 local,只是吐槽他名字起得容易让人产生歧义)

更加细致的解释包括公式、符号意义,参见原文。

2.2 Instance Matching and Aggregation

这里感觉有点类似HTT(Gao et al., AAAI 2019)那篇,不同于之前简单的对每个类别的support set instances 表示取均值,而是利用local aggregation之后的\(\hat{q}\)和每个类别的K个instance即\(\hat{s}_k,k=1,..,K\)做线性非线性变换如公式Eq (11),得到i类中每个句子与q之间的匹配得分,之后Eq (12),attention聚合得到train-support set中每个类别(\(\{s_i^k\}_{k=1}^K,i=1,...,N\))所有instances聚合后的的prototype。

2.3 Class Matching

给query set中的instance归类,class-level matching 函数\(f\),如下公式,即MLP利用prototype \(\hat{s}\)对query分类,得到一个标量--即\(q\)\(\{s^k\}_{k=1}^K\)所属的类别\(i\)的matching score,不同于之前仅考虑q到每个类别prototype的距离作为得分,之后做softmax取最大值,作为预测类别。

2.4 Joint Training with Inconsistency Measurement

就是让 support instance 离 proto 近一点,就引入新的 loss:\(^{[2]}\)

说的有道理,triplet loss 损失函数\(^{[3]}\).

3 实验结果

没细看,直接引用了\(^{[2]}\)

  1. 比 ProtoCNN、ProtoHatt 都好很多
  2. 做了非常多的消融实验:特殊的 pooling 比 avg 或者 max 都好、2.4 的参数共享有用、新 loss 有用、2.1 的一通 cat 操作有用、给 query 归类时 MLP 比单纯算距离好
  3. instance match 画出 attn 的图,发现相似的 query 和 support instance 真的可以对齐

4 总结

基于ProtoCNN,针对query和support set instances以及prototype的表示,引入Local、Instance level Matching and Aggregation,使query与support instances之间有交互得到query、support instances以及prototype的表示,同时最后做分类也不同于传统的prototypical networks用了MLP做分类。

参考

[1] Zhi-Xiu Ye, Zhen-Hua Ling?.Multi-Level Matching and Aggregation Network for Few-Shot Relation Classification.ACL 2019.

[2] 论文笔记 – Multi-Level Matching and Aggregation Network for Few-Shot Relation Classification.https://ivenwang.com/2020/12/11/mlman/.

[3] 洪雨.triplet loss 损失函数.知乎 2020.08.https://zhuanlan.zhihu.com/p/171627918.