ARTICLE · INTELLIGENCE

战地情报 · 详情页

来自尧图项目组的一线实战观察与深度解析

Sentence Transformers 领域自适应(Domain Adaptation)实战指南:从 Adaptive Pre-Training 到 GPL 生成式伪标注

Sentence Transformers 领域自适应(Domain Adaptation)实战指南:从 Adaptive Pre-Training 到 GPL 生成式伪标注 人工智能NLPEmbedding微调【免费下载链接】sentence-transformersState-of-the-Art Embeddings, Retrieval, and Reranking项目地址https://gitcode.com/gh_mirrors/se/sentence-transformers点击查看免费下载导读领域自适应Domain Adaptation的目标是在不依赖人工标注数据的前提下将文本嵌入模型适配到你的特定文本领域。本指南以 examples/sentence_transformer/domain_adaptation/README.md 为核心系统讲解两条主流技术路线——Adaptive Pre-Training自适应预训练含 MLM / TSDAE与 GPLGenerative Pseudo Labeling生成式伪标注并给出论文中的实验数据、MarginMSELoss 的源码级原理以及可直接落地的代码示例。读完本文你将掌握如何用无标签领域语料把通用嵌入模型变成你所在领域的专用模型并理解每条路线的性能收益与计算代价。1. 为什么需要领域自适应Domain Adaptation vs. Unsupervised Learning句子嵌入模型的训练通常依赖标注数据如 Embedding Model Datasets Collection 中的数据集。但当你的语料来自某个特定领域如 AskUbuntu 论坛、法律文书、医学问答、客服对话时通用模型在该领域的检索与匹配效果往往不尽如人意而收集领域内的标注数据又成本高昂。领域自适应正是为了解决这一困境无监督学习Unsupervised Learning仓库在 examples/sentence_transformer/unsupervised_learning/README.md 中汇总了 TSDAE、SimCSE、CT、MLM、GenQ 等方法。它们的共同点是只需要文本本身即可学习语义上有意义的句子嵌入。但正如该文档明确指出的无监督方法在多数情况下性能较差无法真正学习到领域特有的概念尤其在语义搜索任务给定 query 找相关 passage上表现不佳。领域自适应Domain Adaptation更好的思路是——你有一份无标签的领域语料例如 AskUbuntu 的全部帖子标题 一份现有的标注训练集。先利用无标签语料做无监督预训练再在现有标注数据上微调从而把通用知识迁移到你的领域。一句话总结无监督学习只用了领域文本领域自适应 领域文本上的无监督预训练 现有标注数据上的监督微调。2. Adaptive Pre-Training先在领域语料上预训练再在标注集上微调Adaptive Pre-Training 的流程非常直观先用目标语料做无监督预训练如 MLM 或 TSDAE再把预训练好的模型在现有训练数据集上继续微调2.1 两种核心预训练任务MLMMasked Language Model即 BERT 的预训练方式——随机遮盖输入 token让模型预测被遮住的词。仓库提供了开箱即用的脚本 train_mlm.py# 仅提供训练语料 python train_mlm.py distilbert-base path/train.txt # 额外提供验证语料可选 python train_mlm.py distilbert-base path/train.txt path/dev.txttrain.txt / dev.txt 中每一行被视为 Transformer 网络的一条输入即一个句子或段落。注意仅运行 MLM 不会得到好的句子嵌入正确的用法是先在你的领域数据上继续 MLM 预训练再用已有标注数据如 NLI、Paraphrases、STS见 examples/sentence_transformer/training 下的示例做监督微调。TSDAETransformer-based Sequential Denoising AutoEncoder训练时编码器把损坏的句子论文中约删除 60% 的词编码为定长向量解码器则尝试从该句子嵌入重建原始句子为了高质量重建编码器必须把语义完整捕获到句子嵌入中。推理时只使用编码器生成嵌入。仓库在 TSDAE README 中给出了完整训练代码核心是DenoisingAutoEncoderLossimport random from datasets import Dataset from sentence_transformers import SentenceTransformer from sentence_transformers.sentence_transformer.losses import DenoisingAutoEncoderLoss from sentence_transformers.trainer import SentenceTransformerTrainer from sentence_transformers.training_args import SentenceTransformerTrainingArguments # 1. 定义 SentenceTransformer 模型 model SentenceTransformer(google-bert/bert-base-uncased) # 2. 一些示例句子 sentences [ This is an example sentence., Each sentence will be noised and reconstructed., TSDAE learns good sentence embeddings., Sentence Transformers make it easy to train models., ] dataset Dataset.from_dict({text: sentences}) def noise_transform(batch, del_ratio0.6): noisy [] for text in batch[text]: words text.split() keep_prob 1.0 - del_ratio kept_words [w for w in words if random.random() keep_prob] noisy.append( .join(kept_words)) return {noisy: noisy, text: batch[text]} # 3. 添加懒变换在训练时即时为句子加噪 dataset.set_transform(transformlambda batch: noise_transform(batch), columns[text], output_all_columnsTrue) # 4. 定义 TSDAE 损失 train_loss DenoisingAutoEncoderLoss( model, decoder_name_or_pathgoogle-bert/bert-base-uncased, tie_encoder_decoderTrue, ) # 5. 初始化训练参数与 Trainer args SentenceTransformerTrainingArguments( output_diroutput/tsdae-example, num_train_epochs1, per_device_train_batch_size4, ) trainer SentenceTransformerTrainer( modelmodel, argsargs, train_datasetdataset, losstrain_loss, ) # 6. 训练并保存模型 trainer.train() model.save_pretrained(output/tsdae-example/final)从源码 denoising_auto_encoder.py 可以确认其底层机制解码器由AutoModelForCausalLM加载配置为is_decoderTrue、add_cross_attentionTrue因此解码器必须包含XXXLMHead类如 BertLMHeadtie_encoder_decoderTrue默认时编码器与解码器共享权重_tie_encoder_decoder_weights将解码器参数绑定到编码器既提升性能又显著减少显存占用要求编码器与解码器架构一致forward中编码器产出sentence_embedding作为encoder_hidden_states送入解码器形状(bsz, hdim) - (bsz, 1, hdim)解码器以原始句子去掉最后一个 token为输入、原始句子去掉第一个 token为标签用CrossEntropyLoss计算语言建模损失。2.2 论文实验数据预训练带来多大提升在 TSDAE 论文中作者在4 个领域特定的句子嵌入任务上评估了多种领域自适应方法ApproachAskUbuntuCQADupStackTwitterSciDocsAvgZero-Shot Model54.512.972.269.452.3TSDAE59.414.474.577.656.5MLM60.614.371.876.955.9CT56.413.472.469.753.0SimCSE56.213.171.468.952.4可以看到先在领域语料上预训练再在标注数据上微调相比 Zero-Shot 平均提升最高约 8 个点。在 GPL 论文中同样的方法被用于语义搜索给定短查询找到相关段落提升最高可达 10 个点ApproachFiQASciFactBioASQTREC-COVIDCQADupStackRobust04AvgZero-Shot Model26.757.152.966.129.639.045.2TSDAE29.362.855.576.131.839.449.2MLM30.260.051.369.530.438.846.7ICT27.058.355.369.731.337.446.5SimCSE26.755.053.268.329.037.945.0CD27.062.747.765.430.634.544.7CT28.355.649.963.830.535.944.02.3 Adaptive Pre-Training 的代价Adaptive Pre-Training 有一个明显缺点计算开销高。你必须先在领域语料上跑一轮无监督预训练再在标注训练集上跑一轮监督微调而标注训练集可能相当庞大例如all-*-v1系列模型是在超过 10 亿训练对上训练的。这意味着两条流水线都需要完整的训练时间和 GPU 资源。3. GPLGenerative Pseudo Labeling生成式伪标注GPL如all-mpnet-base-v2然后把它适配到你的特定领域无需从零预训练训练时间越长模型效果越好。论文实验中作者在单张 V100-GPU 上训练约 1 天。GPL 还可以与 Adaptive Pre-Training 叠加使用例如先 TSDAE 再 GPL获得进一步的性能提升。3.1 GPL 的三步流程GPL 分三个阶段工作第一步Query Generation查询生成对于领域语料中的一段文本先用一个 T5 模型为该文本生成可能的查询。例如文本是Python is a high-level general-purpose programming language模型可能生成查询What is Python。仓库在 GenQ 教程 中给出了具体实现from transformers import T5Tokenizer, T5ForConditionalGeneration import torch tokenizer T5Tokenizer.from_pretrained(BeIR/query-gen-msmarco-t5-large-v1) model T5ForConditionalGeneration.from_pretrained(BeIR/query-gen-msmarco-t5-large-v1) model.eval() para Python is an interpreted, high-level and general-purpose programming language. Pythons design philosophy emphasizes code readability with its notable use of significant whitespace. Its language constructs and object-oriented approach aim to help programmers write clear, logical code for small and large-scale projects. input_ids tokenizer.encode(para, return_tensorspt) with torch.no_grad(): outputs model.generate( input_idsinput_ids, max_length64, do_sampleTrue, top_p0.95, num_return_sequences3, ) print(Paragraph:) print(para) print(\nGenerated Queries:) for i in range(len(outputs)): query tokenizer.decode(outputs[i], skip_special_tokensTrue) print(f{i 1}: {query})这里使用 Top-p (nucleus) sampling 采样因此每次会生成不同的查询。前身方法GenQ来自 BEIR 论文只做到这一步把生成查询, 段落当作正样本对用MultipleNegativesRankingLoss训练 Bi-Encoder。GPL 则是 GenQ 的改进版。第二步Negative Mining负样本挖掘针对生成的查询What is Python从语料中挖掘负样本段落——即与查询相似、但用户不会认为相关的段落。例如Java is a high-level, class-based, object-oriented programming language.就是这样一个负样本。挖掘采用稠密检索使用现有的文本嵌入模型检索与给定查询相关的段落取回但不直接当作标签。第三步Pseudo Labeling伪标注问题在于负样本挖掘可能挖到实际上与查询相关的段落比如另一段对What is Python的定义。为解决此问题GPL 使用一个 Cross-Encoder 对所有 (query, passage) 对打分。Cross-Encoder 与 Bi-Encoder 的区别在于Bi-Encoder 分别编码两句话得到嵌入 u、v再用余弦相似度比较可索引、可快速检索而 Cross-Encoder 把两个句子同时送入Transformer直接输出一个 0~1 之间的相似度分数精度更高但不产生句子嵌入无法用于大规模索引。在 GPL 中恰好需要精确的逐对打分因此 Cross-Encoder 是伪标注的理想工具。用法见 cross_encoder_usage.pyfrom sentence_transformers.cross_encoder import CrossEncoder model CrossEncoder(cross-encoder/ms-marco-MiniLM-L6-v2) scores model.predict([[My first, sentence pair], [Second text, pair]])第四步Training训练得到三元组(生成查询, 正样本段落, 挖掘出的负样本段落)以及 Cross-Encoder 对(query, positive)和(query, negative)的打分后就可以用 MarginMSELoss 训练文本嵌入模型。伪标注这一步至关重要它正是 GPL 相比前身方法 QGen 性能提升的来源QGen 简单地把段落当作正样本1或负样本0而 GPL 借助 MarginMSELoss Cross-Encoder 识别出“部分相关”或“高度相关”的段落并教会嵌入模型这些段落对于给定查询也是相关的例如对于生成查询what is futures contract负样本挖掘取回的段落中有一部分与查询部分相关或高度相关硬性当作负样本会误导模型Cross-Encoder 给出的软分数则保留了这种相关性梯度。3.2 GPL 在语义搜索上的实验对比下表给出 GPL 与 Adaptive Pre-TrainingMLM、TSDAE的对比可见GPL 可以叠加在 TSDAE 等预训练之上获得最高平均分ApproachFiQASciFactBioASQTREC-COVIDCQADupStackRobust04AvgZero-Shot model26.757.152.966.129.639.045.2TSDAE GPL33.367.362.874.035.142.152.4GPL33.165.261.671.734.442.151.4TSDAE29.362.855.576.131.839.449.2MLM30.260.051.369.530.438.846.74. MarginMSELoss 源码级原理GPL 训练的引擎GPL 训练的最后一环是 MarginMSELoss。该损失的数学定义是计算预测的边界sim(Query, Pos) - sim(Query, Neg)与金标准边界gold_sim(Query, Pos) - gold_sim(Query, Neg)之间的 MSE。默认sim()为点积gold_sim通常来自教师模型在 GPL 中即 Cross-Encoder 的软分数。源码 margin_mse.py 的关键实现要点输入格式(query, document_one, document_two)三元组或(query, positive, negative_1, ..., negative_n)多负样本形式标签可以是“正负分数之差”长度 负样本数也可以是“正样本分数 各负样本分数”的列表长度 负样本数 1后者会在forward中自动转换为差值labels[:, 0].unsqueeze(1) - labels[:, 1:]与 MultipleNegativesRankingLoss 的本质区别后者假定两个文档严格一正一负而 MarginMSELoss 允许两个文档都相关或都不相关只要求保留“哪个更相关”的相对顺序。这正好契合 GPL 场景——负样本挖掘出的段落可能是部分相关的代价同一批 64 的 batch 中MultipleNegativesRankingLoss 会把一个 query 与 128 个文档比较而 MarginMSELoss 一个 query 只与 2 个文档比较训练速度慢得多使用多个负样本会更慢支持知识蒸馏的多种标签形式既可用带硬分数的数据集也可用教师模型similarity_pairwise(emb_q, emb_p1) - similarity_pairwise(emb_q, emb_p2)现场计算软标签还支持多负样本蒸馏——这与 GPL 中“用 Cross-Encoder 打分的 (query, passage) 对作为蒸馏标签”的模式完全一致。5. 如何选择与组合决策小结综合以上分析两条路线可以这样选维度Adaptive Pre-TrainingMLM / TSDAEGPL是否需要标注数据需要预训练后仍需在标注集上微调不需要在已微调模型上直接应用训练语料要求无标签领域文本一行一个句子/段落领域文档语料即可计算开销高预训练 微调两条流水线相对可控在现成微调模型上训练典型场景领域句子嵌入 / 语义相似度领域语义搜索 / 稠密检索组合方式可先 TSDAE/MLM 再叠加 GPL获得最优效果可与 Adaptive Pre-Training 叠加仓库对这条路的完整脉络也有清晰说明在 examples/sentence_transformer/unsupervised_learning/README.md 中GenQ 一节明确写着“本方法已在 GPL 中被改进见 Domain Adaptation”而 MarginMSELoss 的 docstring 也直接引用了Unsupervised Learning Domain Adaptation作为参考文档。三者互为印证构成了完整的“无监督预训练 → 查询生成 → 负样本挖掘 → 伪标注训练”知识体系。6. GPL 代码获取与使用GPL 的官方代码在 UKPLab 的 gpl 仓库中原文档给出的地址为 https://github.com/UKPLab/gpl。设计目标是开箱即用你只需要传入自己的语料库其余步骤查询生成、负样本挖掘、伪标注、训练均由训练代码自动处理。若要亲自动手实践其中的组件可以按以下路径组合仓库资源用 MLM 脚本 或 TSDAE 脚本 在领域语料上做预训练用 GenQ 查询生成示例 中的 T5 模型生成查询用 Cross-Encoder 对 (query, passage) 对打分获得软标签用 MarginMSELoss SentenceTransformerTrainer训练嵌入模型。7. 引用与致谢如果本指南对你有所帮助欢迎引用以下两篇论文。TSDAE: Using Transformer-based Sequential Denoising Auto-Encoder for Unsupervised Sentence Embedding Learninginproceedings{wang-2021-TSDAE, title TSDAE: Using Transformer-based Sequential Denoising Auto-Encoderfor Unsupervised Sentence Embedding Learning, author Wang, Kexin and Reimers, Nils and Gurevych, Iryna, booktitle Findings of the Association for Computational Linguistics: EMNLP 2021, month nov, year 2021, address Punta Cana, Dominican Republic, publisher Association for Computational Linguistics, pages 671--688, url https://arxiv.org/abs/2104.06979, }GPL: Generative Pseudo Labeling for Unsupervised Domain Adaptation of Dense Retrievalinproceedings{wang-2021-GPL, title GPL: Generative Pseudo Labeling for Unsupervised Domain Adaptation of Dense Retrieval, author Wang, Kexin and Thakur, Nandan and Reimers, Nils and Gurevych, Iryna, journal arXiv preprint arXiv:2112.07577, month 12, year 2021, url https://arxiv.org/abs/2112.07577, }8. 延伸阅读无监督学习方法总览TSDAE / SimCSE / CT / MLM / GenQ / GPL 的横向对比TSDAE 完整训练示例含 AskUbuntu 实验与 MAP 结果MLM 预训练脚本Cross-Encoder 使用指南Bi-Encoder 与 Cross-Encoder 的选型与组合预训练模型列表GPL 可以直接适配的现成微调模型MarginMSELoss 参考文档赞分享人工智能NLPEmbedding微调【免费下载链接】sentence-transformersState-of-the-Art Embeddings, Retrieval, and Reranking项目地址https://gitcode.com/gh_mirrors/se/sentence-transformers点击查看免费下载相关推荐领域自适应终极指南awesome-domain-adaptation项目深度解析与实战应用 领域自适应终极指南awesome domain adaptation项目深度解析与实战应用 领域自适应作为机器学习中解决领域偏移问题的关键技术正在迁移学习机器学习文档从理论到实践Awesome-Domain-Adaptation跨域适应的终极指南从理论到实践Awesome Domain Adaptation跨域适应的终极指南 在人工智能快速发展的今天 跨域适应Domain Adaptation迁移学习机器学习文档NeMo ASR Adapters 实战指南领域适配与多任务微调Domain Adaptation Multi-Task Fine-tuningNeMo ASR Adapters 实战指南领域适配与多任务微调Domain Adaptation Multi Task Fine tuning 导读人工智能语音音频大模型深度学习上一篇如何通过Kyverno与Harbor集成实现镜像仓库安全策略全实践下一篇华硕笔记本温度与性能管理全流程G-Helper 终极使用教程创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
RELATED READING

延伸阅读

更多一线实战笔记与深度复盘,助您持续精进