【论文笔记】SDCL: Self-Distillation Contrastive Learning for Chinese Spell Checking

news/2024/4/17 7:57:22

文章目录

  • 论文信息
  • Abstract
  • 1. Introduction
  • 2. Methodology
  • 2.1 The Main Model
    • 2.2 Contrastive Loss
    • 2.3 Implementation Details(Hyperparameters)
  • 3. Experiments
  • 代码实现
  • 个人总结
    • 值得借鉴的地方

论文信息

论文地址:https://arxiv.org/pdf/2210.17168.pdf

Abstract

论文提出了一种token-level的自蒸馏对比学习(self-distillation contrastive learning)方法。

1. Introduction

在这里插入图片描述

传统方法使用BERT后,会对confusion chars进行聚类,但使用作者提出的方法,会让其变得分布更均匀。

confusion chars: 指的应该是易出错的字。

2. Methodology

2.1 The Main Model

作者提取特征的方式:① 先用MacBERT得到hidden states,然后用word embedding和hidden states进行点乘。写成公式为:

H = M a c B E R T ( X ) ⋅ W \bf{H} = MacBERT(X) \cdot W H=MacBERT(X)W

这里的 W W W 应该就是BERT最前面的embedding层对X编码后的向量。

后面就是正常接个输出层再计算CrossEntropyLoss

2.2 Contrastive Loss

在这里插入图片描述

基本思路:让错字token的特征向量和其对应正确字的token特征向量距离越近越好。这样BERT就能拿着错字,然后编码出对应正确字的向量,最后的预测层就能预测对了。

作者的做法:

  1. 错误句子从左边进入BERT,正确句子从右边进入BERT
  2. 对于错字,进行对比学习,让其与对应的正确字的特征向量距离越近越好。即这个错字的正样本为
  3. 将错误句子的其他token作为错字的负样本,使错字token的特征向量与其他向量的距离越远越好。上图中,字有5个负样本,即我、有、吃、旱、饭

上图中双头实线(↔)表示这两个token要距离越近越好,双头虚线表示这两个token要距离越远越好

损失函数公式如下:

L c = − ∑ i = 1 n L ( x ~ i ) log ⁡ exp ⁡ ( sim ⁡ ( h ~ i , h i ) / τ ) ∑ j = 1 n exp ⁡ ( sim ⁡ ( h ~ i , h j ) / τ ) L_c = -\sum_{i=1}^n \Bbb{L}\left(\tilde{x}_i\right) \log \frac{\exp \left(\operatorname{sim}\left(\tilde{h}_i, h_i\right) / \tau\right)}{\sum_{j=1}^n \exp \left(\operatorname{sim}\left(\tilde{h}_i, h_j\right) / \tau\right)} Lc=i=1nL(x~i)logj=1nexp(sim(h~i,hj)/τ)exp(sim(h~i,hi)/τ)

其中:

  • n n n : 为n个token
  • L ( x ~ i ) \Bbb{L}\left(\tilde{x}_i\right) L(x~i): 当 x i x_i xi为错字时, L ( x ~ i ) = 1 \Bbb{L}\left(\tilde{x}_i\right)=1 L(x~i)=1,否则为 0 0 0。即只算错字的损失
  • sim ( ⋅ ) \text{sim}(\cdot) sim():余弦相似度函数
  • h ~ i \tilde{h}_i h~i: 正确句子(右边BERT)输出的token的特征向量
  • h i h_i hi:错误句子(左边BERT)输出的token的特征向量
  • τ \tau τ:温度超参

上面损失使用CrossEntropyLoss实现。

作者还为右边的BERT增加了一个Loss L y L_y Ly,目的是让右边可以输出它的输入,即copy-paste任务。

最终的损失如下:

L = L x + α L y + β L c L = L_x+\alpha L_y+\beta L_c L=Lx+αLy+βLc

2.3 Implementation Details(Hyperparameters)

  • BERT:MacBERT
  • optimizer: AdamW
  • 学习率: 7e-5
  • batch_size: 48
  • λ \lambda λ: 0.9 (TODO,作者说的这个lambda不知道是啥)
  • α \alpha α: 1
  • β \beta β: 0.5
  • τ \tau τ: 0.9
  • epoch: 20次

3. Experiments

在这里插入图片描述

代码实现

import torch
import torch.nn as nn
from transformers import BertTokenizerFast, BertForMaskedLM
import torch.nn.functional as Fclass SDCLModel(nn.Module):def __init__(self):super(SDCLModel, self).__init__()self.tokenizer = BertTokenizerFast.from_pretrained('hfl/chinese-macbert-base')self.model = BertForMaskedLM.from_pretrained('hfl/chinese-macbert-base')self.alpha = 1self.beta = 0.5self.temperature = 0.9def forward(self, inputs, targets=None):"""inputs: 为tokenizer对原文本编码后的输入,包括input_ids, attention_mask等targets:与inputs相同,只不过是对目标文本编码后的结果。"""if targets is not None:# 提取labels的input_idstext_labels = targets['input_ids'].clone()text_labels[text_labels == 0] = -100  # -100计算损失时会忽略else:text_labels = Noneword_embeddings = self.model.bert.embeddings.word_embeddings(inputs['input_ids'])hidden_states = self.model.bert(**inputs).last_hidden_statelogits = self.model.cls(hidden_states * word_embeddings)if targets:loss = F.cross_entropy(logits.view(logits.shape[0] * logits.shape[1], logits.shape[2]), text_labels.view(-1))else:loss = 0.return logits, hidden_states, lossdef extract_outputs(self, outputs):logits, _, _ = outputsreturn logits.argmax(-1)def compute_loss(self, outputs, targets, inputs, detect_targets, *args, **kwargs):logits_x, hidden_states_x, loss_x = outputslogits_y, hidden_states_y, loss_y = self.forward(targets, targets)# FIXMEanchor_samples = hidden_states_x[detect_targets.bool()]positive_samples = hidden_states_y[detect_targets.bool()]negative_samples = hidden_states_x[~detect_targets.bool() & inputs['attention_mask'].bool()]# 错字和对应正确的字计算余弦相似度positive_sim = F.cosine_similarity(anchor_samples, positive_samples)# 错字与所有batch内的所有其他字计算余弦相似度# (FIXME,这里与原论文不一致,原论文说的是与当前句子的其他字计算,但我除了for循环,不知道该怎么写)negative_sim = F.cosine_similarity(anchor_samples.unsqueeze(1), negative_samples.unsqueeze(0), dim=-1)sims = torch.concat([positive_sim.unsqueeze(1), negative_sim], dim=1) / self.temperaturesim_labels = torch.zeros(sims.shape[0]).long().to(self.args.device)loss_c = F.cross_entropy(sims, sim_labels)self.loss_c = float(loss_c)  # 记录一下return loss_x + self.alpha * loss_y + self.beta * loss_cdef get_optimizer(self):return torch.optim.AdamW(self.parameters(), lr=7e-5)def predict(self, src):src = ' '.join(src.replace(" ", ""))inputs = self.tokenizer(src, return_tensors='pt').to(self.args.device)outputs = self.forward(inputs)outputs = self.extract_outputs(outputs)[0][1:-1]return self.tokenizer.decode(outputs).replace(' ', '')

个人总结

值得借鉴的地方

  1. 作者并没有直接使用BERT的输出作为token embedding,而是使用点乘的方式融合了BERT的输出和word embeddings

https://www.xjx100.cn/news/3118790.html

相关文章

ElasticSearch之cat indices API

命令样例如下: curl -X GET "https://localhost:9200/_cat/indices?vtrue&pretty" --cacert $ES_HOME/config/certs/http_ca.crt -u "elastic:ohCxPHQBEs5*lo7F9"执行结果输出如下: health status index uuid …

PHP项目用docker一键部署

公司新项目依赖较多,扩展版本参差不一,搭建环境复杂缓慢,所以搭建了一键部署的功能。 docker-compose build 构建docker docker-compose up 更新docker docker-compose up -d 后台运行docker docker exec -it docker-php-1 /bin/bas…

【UE】中文字体 发光描边材质

效果 步骤 1. 先将我们电脑中存放在“C:\Windows\Fonts”路径下的字体导入UE 点击“全部选是” 导入成功后如下 2. 打开导入的“SIMSUN_Font”,将字体缓存类型设置为“离线” 点击“是” 这里我选择:宋体-常规-20 展开细节面板中的导入选项 勾选“使用距…

【从JVM看Java,三问继承和多态,是什么?为什么?怎么做?深度剖析JVM的工作原理】

系列文章: 《计算机底层原理专栏》:欢迎大家订阅学习,能够帮助到各位就是对我最大的鼓励! 文章目录 系列文章目录前言一、JVM是什么二、什么是继承三、什么是多态总结 前言 这篇文章聚焦JVM的实现原理,我更专注于从一…

Flink-时间窗口

在流数据处理应用中,一个很重要、也很常见的操作就是窗口计算。所谓的“窗口”,一 般就是划定的一段时间范围,也就是“时间窗”;对在这范围内的数据进行处理,就是所谓的 窗口计算。所以窗口和时间往往是分不开的。 时…

k8s中Pod控制器简介,ReplicaSet、Deployment、HPA三种处理无状态pod应用的控制器介绍

目录 一.Pod控制器简介 二.ReplicaSet(简写rs) 1.简介 (1)主要功能 (2)rs较完整参数解释 2.创建和删除 (1)创建 (2)删除 3.扩容和缩容 &#xff08…

ubuntu下训练自己的yolov5数据集

参考文档 yolov5-github yolov5-github-训练文档 csdn训练博客 一、配置环境 1.1 安装依赖包 前往清华源官方地址 选择适合自己的版本替换自己的源 # 备份源文件 sudo cp /etc/apt/sources.list /etc/apt/sources.list_bak # 修改源文件 # 更新 sudo apt update &&a…

企业微信应用文本消息

应用支持推送文本、图片、视频、文件、图文等类型,本篇主要实现发送应用文本消息。 获取企业凭证 发送应用消息首先需要获取调用凭证access_token,此处的凭证为企业凭证,可通过企业授权安装时返回的授权信息中获取access_token;之…