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

文章目录

  • 论文信息
  • 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

本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若转载,请注明出处:http://www.mzph.cn/news/183147.shtml

如若内容造成侵权/违法违规/事实不符,请联系多彩编程网进行投诉反馈email:809451989@qq.com,一经查实,立即删除!

相关文章

idea doc 注释 插件及使用

开启rendered view https://blog.csdn.net/Leiyi_Ann/article/details/124145492 生成doc https://blog.csdn.net/qq_42581682/article/details/105018239 把注释加到类名旁边插件 https://blog.csdn.net/qq_30231473/article/details/128825306

聚类分析例题 (多元统计分析期末复习)

例一 动态聚类,K-means法,随机选取凝聚点(题目直接给出) 已知5个样品的观测值为:1,4,5,7,11。试用K均值法分为两类(凝聚点分别取1,4与1,11) 解&…

找不到 sun.misc.BASE64Decoder ,sun.misc.BASE64Encoder 类

找不到 sun.misc.BASE64Decoder ,sun.misc.BASE64Encoder 类 1. 现象 idea 引用报错 找不到对应的包 import sun.misc.BASE64Decoder; import sun.misc.BASE64Encoder;2. 原因 因为sun.misc.BASE64Decoder和sun.misc.BASE64Encoder是Java的内部API,通…

VR虚拟教育展厅,为教学领域开启创新之路

线上虚拟展厅是一项全新的展示技术,可以为参展者带来不一样的观展体验。传统的实体展览存在着空间限制、时间限制以及高昂的成本,因此对于教育领域来说,线上虚拟教育展厅的出现,可以对传统教育方式带来改革,凭借强大的…

ORA-00837: Specified value of MEMORY_TARGET greater than MEMORY_MAX_TARGET

有个11g rac环境,停电维护后,orcl1正常启动了,orcl2启动报错如下 SQL*Plus: Release 11.2.0.4.0 Production on Wed Nov 29 14:04:21 2023 Copyright (c) 1982, 2013, Oracle. All rights reserved. Connected to an idle instance. SYS…

从0开始学习JavaScript--JavaScript 模板字符串的全面应用

JavaScript 模板字符串是 ES6 引入的一项强大特性,它提供了一种更优雅、更灵活的字符串拼接方式。在本文中,将深入探讨模板字符串的基本语法、高级用法以及在实际项目中的广泛应用,通过丰富的示例代码带你领略模板字符串的魅力。 模板字符串…

亚马逊云科技基于 Polygon 推出首款 Amazon Managed Blockchain Access,助 Web3 开发人员降低区块链节点运行成本

2023 年 11 月 26 日,亚马逊 (Amazon) 旗下 Amazon Web Services(Amazon)在其官方博客上宣布,Amazon Managed Blockchain (AMB) Access 已支持 Polygon Proof-of-Stake(POS) 网络,并将满足各种场景的需求,包…

删除list中除最后一个之外所有的数据

1.你可以新建一个list List<Integer> listnew ArrayList<>();int i0;while (i<100){list.add(i);}List<Integer> subList list.subList(list.size()-1, list.size());System.out.println("原list大小--"list.size());System.out.println("…

群晖安装portainer

一、下载镜像 打开【Container Manager】 ,搜索portainer&#xff0c;双击【6053537/portainer-ce】下载汉化版本 二、创建映射文件夹 打开【File Station】&#xff0c;在docker目录下创建【portainer】文件夹 三、开启SSH 群晖 - 【控制面板】-【终端机和SNMP】 勾选【启动…

第二十章 多线程总结

继承Thread 类 Thread 类时 java.lang 包中的一个类&#xff0c;从类中实例化的对象代表线程&#xff0c;程序员启动一个新线程需要建立 Thread 实例。 Thread 对象需要一个任务来执行&#xff0c;任务是指线程在启动时执行的工作&#xff0c;start() 方法启动线程&…

五、初识FreeRTOS之FreeRTOS的任务创建和删除

本节主要学习以下内容&#xff1a; 1&#xff0c;任务创建和删除的API函数&#xff08;熟悉&#xff09; 2&#xff0c;任务创建和删除&#xff08;动态方法&#xff09;&#xff08;掌握&#xff09; 3&#xff0c;任务创建和删除&#xff08;静态方法&#xff09;&#xf…

mongodb基本操作命令

mongodb快速搭建及使用 1.mongodb安装1.1 docker安装启动mongodb 2.mongo shell常用命令2.1 插入文档2.1.1 插入单个文档2.1.2 插入多个文档2.1.3 用脚本批量插入 2.2 查询文档 前言&#xff1a;本篇默认你是对nongodb的基础概念有了了解&#xff0c;操作是非常基础的。但是与关…

微信小程序——给按钮添加点击音效

今天来讲解一下如何给微信小程序的按钮添加点击音效 注意&#xff1a;这里的按钮不一定只是 <button>&#xff0c;也可以是一张图片&#xff0c;其实只是添加一个监听点击事件的函数而已 首先来看下按钮的定义 <button bind:tap"onInput" >点我有音效&…

C++面向对象复习笔记暨备忘录

C指针 指针作为形参 交换两个实际参数的值 #include <iostream> #include<cassert> using namespace std;int swap(int *x, int* y) {int a;a *x;*x *y;*y a;return 0; } int main() {int a 1;int b 2;swap(&a, &b);cout << a << &quo…

【开源视频联动物联网平台】为什么需要物联网网关?

在一些物联网项目中&#xff0c;物联网网关这一产品经常被涉及。那么&#xff0c;物联网网关究竟有何作用&#xff1f;具备哪些功能&#xff1f;同时&#xff0c;我们也发现有些物联网设备并不需要网关。那么&#xff0c;究竟在何时需要物联网网关呢&#xff1f; 物联网的架构…

LaTeX插入裁剪后的pdf图像

画图 VSCode Draw.io Integration插件 有数学公式的打开下面的选项&#xff1a; 导出 File -> Export -> .svg导出成svg格式的文件。然后用浏览器打开svg文件后CtrlP选择另存为PDF&#xff0c;将图片存成pdf格式。 裁剪 只要安装了TeXLive&#xff0c;就只需要在图…

LVS-NAT实验

实验前准备&#xff1a; LVS负载调度器&#xff1a;ens33&#xff1a;192.168.20.11 ens34&#xff1a;192.168.188.3 Web1节点服务器1&#xff1a;192.168.20.12 Web2节点服务器2&#xff1a;192.168.20.13 NFS服务器&#xff1a;192.168.20.14 客户端&#xff08;win11…

ESD静电试验方法及标准

文章目录 概述静电放电抗扰标准静电放电实验室的型式试验静电放电试验配置静电放电试验方法 静电放电等级 参考静电放电发生器&#xff08;ESD&#xff09;试验方法及标准 概述 在低湿度环境下通过摩擦使人体充电的人体在与设备接触时可能会放电&#xff0c;静电放电的后果是&…

uniapp 打包的 IOS打开白屏 uniapp打包页面空白

uniapp的路由跟vue一样,有hash模式和history模式, 使用 URL 的 hash 来模拟一个完整的 URL,于是当 URL 改变时,页面不会重新加载。 如果不想要很丑的 hash,我们可以用路由的 history 模式,这种模式充分利用 history.pushState API 来完成 URL 跳转而无须重新加载页面。…

PHP项目用docker一键部署

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