- 人工智能
- NLP
- Embedding
- 微调
【免费下载链接】sentence-transformers
State-of-the-Art Embeddings, Retrieval, and Reranking
导读
领域自适应(Domain Adaptation)的目标是:在不依赖人工标注数据的前提下,将文本嵌入模型适配到你的特定文本领域。本指南以 examples/sentence_transformer/domain_adaptation/README.md 为核心,系统讲解两条主流技术路线——Adaptive Pre-Training(自适应预训练,含 MLM / TSDAE)与 GPL(Generative 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 两种核心预训练任务
MLM(Masked 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 下的示例)做监督微调。
TSDAE(Transformer-based Sequential Denoising AutoEncoder):训练时,编码器把损坏的句子(论文中约删除 60% 的词)编码为定长向量,解码器则尝试从该句子嵌入重建原始句子;为了高质量重建,编码器必须把语义完整捕获到句子嵌入中。推理时只使用编码器生成嵌入。仓库在 TSDAE README 中给出了完整训练代码,核心是DenoisingAutoEncoderLoss:
import 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_ratio=0.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(transform=lambda batch: noise_transform(batch), columns=["text"], output_all_columns=True) # 4. 定义 TSDAE 损失 train_loss = DenoisingAutoEncoderLoss( model, decoder_name_or_path="google-bert/bert-base-uncased", tie_encoder_decoder=True, ) # 5. 初始化训练参数与 Trainer args = SentenceTransformerTrainingArguments( output_dir="output/tsdae-example", num_train_epochs=1, per_device_train_batch_size=4, ) trainer = SentenceTransformerTrainer( model=model, args=args, train_dataset=dataset, loss=train_loss, ) # 6. 训练并保存模型 trainer.train() model.save_pretrained("output/tsdae-example/final")从源码 denoising_auto_encoder.py 可以确认其底层机制:
- 解码器由
AutoModelForCausalLM加载,配置为is_decoder=True、add_cross_attention=True,因此解码器必须包含XXXLMHead类(如 BertLMHead); tie_encoder_decoder=True(默认)时,编码器与解码器共享权重,_tie_encoder_decoder_weights将解码器参数绑定到编码器,既提升性能又显著减少显存占用;要求编码器与解码器架构一致;forward中,编码器产出sentence_embedding,作为encoder_hidden_states送入解码器(形状(bsz, hdim) -> (bsz, 1, hdim)),解码器以原始句子(去掉最后一个 token)为输入、原始句子(去掉第一个 token)为标签,用CrossEntropyLoss计算语言建模损失。
2.2 论文实验数据:预训练带来多大提升?
在 TSDAE 论文中,作者在4 个领域特定的句子嵌入任务上评估了多种领域自适应方法:
| Approach | AskUbuntu | CQADupStack | SciDocs | Avg | |
|---|---|---|---|---|---|
| Zero-Shot Model | 54.5 | 12.9 | 72.2 | 69.4 | 52.3 |
| TSDAE | 59.4 | 14.4 | 74.5 | 77.6 | 56.5 |
| MLM | 60.6 | 14.3 | 71.8 | 76.9 | 55.9 |
| CT | 56.4 | 13.4 | 72.4 | 69.7 | 53.0 |
| SimCSE | 56.2 | 13.1 | 71.4 | 68.9 | 52.4 |
可以看到,先在领域语料上预训练再在标注数据上微调,相比 Zero-Shot 平均提升最高约 8 个点。
在 GPL 论文中,同样的方法被用于语义搜索(给定短查询找到相关段落),提升最高可达 10 个点:
| Approach | FiQA | SciFact | BioASQ | TREC-COVID | CQADupStack | Robust04 | Avg |
|---|---|---|---|---|---|---|---|
| Zero-Shot Model | 26.7 | 57.1 | 52.9 | 66.1 | 29.6 | 39.0 | 45.2 |
| TSDAE | 29.3 | 62.8 | 55.5 | 76.1 | 31.8 | 39.4 | 49.2 |
| MLM | 30.2 | 60.0 | 51.3 | 69.5 | 30.4 | 38.8 | 46.7 |
| ICT | 27.0 | 58.3 | 55.3 | 69.7 | 31.3 | 37.4 | 46.5 |
| SimCSE | 26.7 | 55.0 | 53.2 | 68.3 | 29.0 | 37.9 | 45.0 |
| CD | 27.0 | 62.7 | 47.7 | 65.4 | 30.6 | 34.5 | 44.7 |
| CT | 28.3 | 55.6 | 49.9 | 63.8 | 30.5 | 35.9 | 44.0 |
2.3 Adaptive Pre-Training 的代价
Adaptive Pre-Training 有一个明显缺点:计算开销高。你必须先在领域语料上跑一轮无监督预训练,再在标注训练集上跑一轮监督微调;而标注训练集可能相当庞大(例如all-*-v1系列模型是在超过 10 亿训练对上训练的)。这意味着两条流水线都需要完整的训练时间和 GPU 资源。
3. GPL:Generative 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. Python's 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_tensors="pt") with torch.no_grad(): outputs = model.generate( input_ids=input_ids, max_length=64, do_sample=True, top_p=0.95, num_return_sequences=3, ) print("Paragraph:") print(para) print("\nGenerated Queries:") for i in range(len(outputs)): query = tokenizer.decode(outputs[i], skip_special_tokens=True) 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.py:
from 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-Training(MLM、TSDAE)的对比,可见GPL 可以叠加在 TSDAE 等预训练之上获得最高平均分:
| Approach | FiQA | SciFact | BioASQ | TREC-COVID | CQADupStack | Robust04 | Avg |
|---|---|---|---|---|---|---|---|
| Zero-Shot model | 26.7 | 57.1 | 52.9 | 66.1 | 29.6 | 39.0 | 45.2 |
| TSDAE + GPL | 33.3 | 67.3 | 62.8 | 74.0 | 35.1 | 42.1 | 52.4 |
| GPL | 33.1 | 65.2 | 61.6 | 71.7 | 34.4 | 42.1 | 51.4 |
| TSDAE | 29.3 | 62.8 | 55.5 | 76.1 | 31.8 | 39.4 | 49.2 |
| MLM | 30.2 | 60.0 | 51.3 | 69.5 | 30.4 | 38.8 | 46.7 |
4. 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-Training(MLM / TSDAE) | GPL |
|---|---|---|
| 是否需要标注数据 | 需要(预训练后仍需在标注集上微调) | 不需要(在已微调模型上直接应用) |
| 训练语料要求 | 无标签领域文本(一行一个句子/段落) | 领域文档语料即可 |
| 计算开销 | 高(预训练 + 微调两条流水线) | 相对可控(在现成微调模型上训练) |
| 典型场景 | 领域句子嵌入 / 语义相似度 | 领域语义搜索 / 稠密检索 |
| 组合方式 | 可先 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 Learning
@inproceedings{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 Retrieval
@inproceedings{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 参考文档
- 人工智能
- NLP
- Embedding
- 微调
【免费下载链接】sentence-transformers
State-of-the-Art Embeddings, Retrieval, and Reranking
相关推荐
领域自适应终极指南:awesome-domain-adaptation项目深度解析与实战应用 🚀
领域自适应终极指南:awesome domain adaptation项目深度解析与实战应用 🚀 领域自适应作为机器学习中解决"领域偏移"问题的关键技术,正在
迁移学习机器学习文档从理论到实践:Awesome-Domain-Adaptation跨域适应的终极指南
从理论到实践:Awesome Domain Adaptation跨域适应的终极指南 在人工智能快速发展的今天, 跨域适应(Domain Adaptation)
迁移学习机器学习文档NeMo ASR Adapters 实战指南:领域适配与多任务微调(Domain Adaptation & Multi-Task Fine-tuning)
NeMo ASR Adapters 实战指南:领域适配与多任务微调(Domain Adaptation & Multi Task Fine tuning) 导读
人工智能语音音频大模型深度学习
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考