HyperAIHyperAI

Command Palette

Search for a command to run...

GAN-BERT:利用少量标注样本进行鲁棒文本分类的生成对抗学习

Danilo Croce Giuseppe Castellucci Roberto Basili

摘要

近年来,基于 Transformer 的架构(如 BERT)在许多自然语言处理任务中取得了令人瞩目的成果。然而,大多数采用的基准数据集包含(有时多达数十万的)大量样本。在许多实际场景中,获取高质量的标注数据既昂贵又耗时;相比之下,表征目标任务的无标注样本通常可以较为容易地收集。图像处理领域已提出一种基于半监督生成对抗网络的有前景方法,用以实现半监督学习。本文提出 GAN-BERT,在生成对抗框架下利用无标注数据扩展 BERT 类架构的微调过程。实验结果表明,对标注样本的需求可以大幅降低(最低仅需 50 至 100 个标注样本),同时在多个句子分类任务中仍能取得良好性能。

一句话总结

罗马大学托尔维加塔分校和亚马逊的研究人员提出了 GAN-BERT,该方法在半监督生成对抗框架下利用无标注数据扩展了类似 BERT 的微调,将所需标注样本减少到仅 50–100 个,同时在多个句子分类任务上保持良好性能。

核心贡献

  • GAN-BERT 通过半监督生成对抗设置扩展 BERT 微调,其中生成器产生伪样本,BERT 充当判别器。
  • 该方法降低标注需求,仅使用 50 到 100 个标注样本即可获得良好的句子分类性能,且在少于 200 个标注样本时结果可与全监督设置相媲美。
  • 半监督对抗方案相较于 BERT 持续提升性能,且不增加推理成本,因为生成器仅在训练阶段使用。

引言

基于 Transformer 的模型(如 BERT)在大型标注数据集上微调后能取得很强的 NLP 性能,但当只有几百个标注样本可用时,其准确率会显著下降,尤其是对于类别较多的分类任务。由于人工标注成本高昂,半监督生成对抗网络是一种有吸引力的替代方案,但它们在 NLP 中的应用此前仅限于一种基于核的 GAN,该 GAN 在固定预计算嵌入上运行,而不更新表示空间。作者提出了 GAN-BERT,将 BERT 作为判别器置于半监督 GAN 框架中,利用无标注数据改进表示,并在少于 200 个标注样本的情况下达到可与全监督训练相媲美的结果。

方法

作者利用半监督 GAN(SS-GAN)在 GAN 框架内实现半监督学习。在该设置中,判别器按照 (k+1)(k + 1)(k+1) 类目标进行训练。真实样本被分类到目标类别 (1,,k)(1, \dots, k)(1,,k) 之一,而生成的样本被分类到第 (k+1)(k + 1)(k+1) 类。形式上,令 DDDGGG 分别表示判别器和生成器,pdp_dpdpGp_\mathcal{G}pG 分别表示真实数据分布和生成样本。DDD 的损失函数定义为 LD=LDsup+LDunsupL_D = L_{D_{sup}} + L_{D_{unsup}}LD=LDsup+LDunsup,其中:

LDsup=Ex,ypdlog[pm(y^=yx,y(1,,k))]LDunsup=Expdlog[1pm(y^=yx,y=k+1)]ExGlog[pm(y^=yx,y=k+1)]\begin{array}{c} L_{D_{sup}} = - \mathbb{E}_{x, y \sim p_d} \log [ p_m(\hat{y} = y | x, y \in (1, \dots, k)) ] \\ L_{D_{unsup}} = - \mathbb{E}_{x \sim p_d} \log [ 1 - p_m(\hat{y} = y | x, y = k + 1) ] - \mathbb{E}_{x \sim \mathcal{G}} \log [ p_m(\hat{y} = y | x, y = k + 1) ] \end{array}LDsup=Ex,ypdlog[pm(y^=yx,y(1,,k))]LDunsup=Expdlog[1pm(y^=yx,y=k+1)]ExGlog[pm(y^=yx,y=k+1)]

LDsupL_{D_{sup}}LDsup 衡量在原始 kkk 个类别中为真实样本分配错误类别的误差。LDunsupL_{D_{unsup}}LDunsup 衡量将真实无标注样本错误识别为假样本以及未能识别假样本的误差。生成器损失 LGL_GLG 结合了特征匹配损失和无监督损失,鼓励 GGG 生成中间表示与真实样本相似的样本。

为了将其用于自然语言处理,作者提出了 GAN-BERT,它在微调阶段通过集成 SS-GAN 层来扩展预训练的 BERT 模型。给定输入句子,BERT 生成向量表示,并采用 hCLSh_{CLS}hCLS 表示作为目标任务的句子嵌入。

如下图所示:

该架构在 BERT 之上增加了一个用于分类样本的判别器 DDD 和一个对抗性作用的生成器 GGG。生成器 GGG 是一个多层感知机(MLP),其输入是从 N(μ,σ2)N(\mu, \sigma^2)N(μ,σ2) 中采样的 100 维噪声向量,输出向量 hfakeRdh_{fake} \in \mathbb{R}^dhfakeRd。判别器是另一个 MLP,接收向量 hRdh_* \in \mathbb{R}^dhRd,该向量可以是生成器产生的 hfakeh_{fake}hfake,也可以是来自真实分布的无标注或标注样本的 hCLSh_{CLS}hCLSDDD 的最后一层是 softmax 激活层,输出 k+1k + 1k+1 维 logits 向量。

在训练过程中,系统优化两个相互竞争的损失 LDL_DLDLGL_GLG。在前向步骤中,当采样到真实样本(h=hCLSh_* = h_{CLS}h=hCLS)时,DDD 将其分类到 kkk 个类别之一。当 h=hfakeh_* = h_{fake}h=hfake 时,将其分类到第 k+1k + 1k+1 类。在反向传播过程中,无标注样本仅对 LDunsupL_{D_{unsup}}LDunsup 产生贡献,即只有当它们被错误分类到第 k+1k + 1k+1 类时才会参与损失计算;在其他所有情况下,它们对损失的贡献被屏蔽。标注样本对监督损失 LDsupL_{D_{sup}}LDsup 产生贡献。GGG 生成的样本同时对 LDL_DLDLGL_GLG 产生贡献。更新 DDD 时,BERT 权重也会被修改,以利用标注和无标注数据微调其内部表示。训练结束后,生成器 GGG 被丢弃,保留原始 BERT 模型进行推理,不增加额外计算成本。

实验

实验在主题分类、问题分类、情感分析和自然语言推理任务上,使用逐步增大的标注集和额外的无标注样本,将 GAN-BERT 与微调后的 BERT-base 模型进行比较。当只有少量标注样本可用时,GAN-BERT 相较于 BERT 持续提升性能,而 BERT 在 1% 标签比例下常常发散,且对于类别较多的任务收益更明显。该优势在 SST-5 上同样出现,并在 MNLI 上在约 0.5% 标注样本范围内系统性地存在,此后两个模型表现相似。


用 AI 构建 AI

从创意到上线——通过免费 AI 协同编码、开箱即用的环境和最优惠的 GPU 价格,加速您的 AI 开发。

AI 协同编码
开箱即用的 GPU
最优定价

HyperAI Newsletters

订阅我们的最新资讯
我们会在北京时间 每周一的上午九点 向您的邮箱投递本周内的最新更新
邮件发送服务由 MailChimp 提供