笔记: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]}\)。
- 比 ProtoCNN、ProtoHatt 都好很多
- 做了非常多的消融实验:特殊的 pooling 比 avg 或者 max 都好、2.4 的参数共享有用、新 loss 有用、2.1 的一通 cat 操作有用、给 query 归类时 MLP 比单纯算距离好
- 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.