如何用Huggingface Trainer API基于对比学习微调Sentence Transformer?
对比学习微调Sentence Transformer(MPNET)使用Huggingface Trainer API方案
完全可以用Huggingface Trainer API实现对比学习微调Sentence Transformer类模型(比如MPNET)。Trainer API是通用训练框架,核心是自定义数据处理流程和模型损失计算逻辑,就能适配对比学习这类非分类任务。
具体实现步骤
1. 准备对比学习数据集
对比学习需要构造包含「锚点(Anchor)、正样本(Positive)、负样本(Negative)」的样本对。你可以:
- 使用现成的相似性数据集(如STS-B、MRPC),将相似句子作为正样本,不相似的作为负样本;
- 自行构造:对锚点句子做同义词替换生成正样本,从不同语义类别中选取负样本;
- 采用批量内负采样(无需提前构造负样本):将同批次内的其他样本作为当前样本的负样本,效率更高。
2. 自定义数据处理逻辑
需要将文本转换为模型可接受的token格式,并整理成包含锚点、正/负样本的输入批次:
- 使用
Dataset.map()完成文本tokenization,分别处理锚点、正、负样本; - 自定义
DataCollator,将tokenized后的样本整理成模型forward需要的输入字典(包含锚点、正/负样本的input_ids和attention_mask)。
3. 包装模型并实现对比损失
原Sentence Transformer模型仅输出embedding,需要包装一层来计算对比学习损失(常用InfoNCE或Triplet Loss):
- 继承
PreTrainedModel,以MPNET为backbone; - 在
forward方法中计算锚点、正/负样本的embedding并做归一化; - 实现对比损失计算逻辑,返回包含
loss键的字典供Trainer使用。
4. 配置Trainer并启动训练
关键是设置remove_unused_columns=False(避免Trainer自动删除模型需要的输入列),并传入自定义的模型、数据集、data_collator。
代码示例
from transformers import AutoModel, AutoTokenizer, Trainer, TrainingArguments, PreTrainedModel import torch import torch.nn as nn from torch.nn import functional as F from datasets import Dataset # 1. 自定义对比学习模型 class ContrastiveMPNET(PreTrainedModel): def __init__(self, model_name, temperature=0.1): super().__init__(AutoModel.from_pretrained(model_name).config) self.backbone = AutoModel.from_pretrained(model_name) self.temperature = temperature # 对比学习温度系数,影响损失分布 def forward(self, anchor_inputs, positive_inputs, negative_inputs=None): # 获取<[BOS_never_used_51bce0c785ca2f68081bfa7d91973934]>token的embedding并归一化 def get_embedding(inputs): outputs = self.backbone(**inputs) emb = outputs.last_hidden_state[:, 0, :] return F.normalize(emb, p=2, dim=1) anchor_emb = get_embedding(anchor_inputs) positive_emb = get_embedding(positive_inputs) # 批量内负采样:用同批次其他样本作为负样本 if negative_inputs is None: negative_emb = torch.cat([anchor_emb[1:], anchor_emb[:1]], dim=0) else: negative_emb = get_embedding(negative_inputs) # 计算InfoNCE损失 # 正样本相似度 + 负样本相似度 logits = torch.cat( [torch.matmul(anchor_emb, positive_emb.T).diag().unsqueeze(1), torch.matmul(anchor_emb, negative_emb.T)], dim=1 ) / self.temperature # 标签:正样本对应索引0 labels = torch.zeros(logits.shape[0], dtype=torch.long, device=logits.device) loss = F.cross_entropy(logits, labels) return {"loss": loss} # 2. 准备并处理数据集 sample_data = { "anchor": ["我喜欢吃苹果", "今天天气很好", "机器学习很有趣"], "positive": ["我爱吃苹果", "今日天气很不错", "ML非常有意思"], "negative": ["我讨厌吃香蕉", "今天下雨了", "深度学习很难"] } dataset = Dataset.from_dict(sample_data) tokenizer = AutoTokenizer.from_pretrained("microsoft/mpnet-base") def tokenize_batch(examples): return { "anchor": tokenizer(examples["anchor"], padding="max_length", truncation=True, max_length=64), "positive": tokenizer(examples["positive"], padding="max_length", truncation=True, max_length=64), "negative": tokenizer(examples["negative"], padding="max_length", truncation=True, max_length=64) } tokenized_dataset = dataset.map(tokenize_batch, batched=True) # 自定义DataCollator class ContrastiveCollator: def __call__(self, features): def extract_inputs(feature_key): return { "input_ids": torch.tensor([f[feature_key]["input_ids"] for f in features]), "attention_mask": torch.tensor([f[feature_key]["attention_mask"] for f in features]) } return { "anchor_inputs": extract_inputs("anchor"), "positive_inputs": extract_inputs("positive"), "negative_inputs": extract_inputs("negative") } data_collator = ContrastiveCollator() # 3. 配置Trainer并训练 model = ContrastiveMPNET("microsoft/mpnet-base", temperature=0.1) training_args = TrainingArguments( output_dir="./contrastive_mpnet_checkpoints", per_device_train_batch_size=8, num_train_epochs=3, learning_rate=2e-5, logging_steps=5, save_strategy="epoch", remove_unused_columns=False # 必须设置,否则Trainer会删除模型需要的输入列 ) trainer = Trainer( model=model, args=training_args, train_dataset=tokenized_dataset, data_collator=data_collator ) trainer.train()
关键指导建议
- 损失函数选择:优先用InfoNCE损失(适配批量内负采样,训练效率高);如果有明确的正负样本对,也可以用Triplet Loss(需设置合适的margin)。
- 超参数调整:
- 温度系数:建议在0.05~0.5之间,过小会让损失过于集中,过大则损失分布太平缓;
- 学习率:比分类任务低,建议2e-5~5e-5;
- 批次大小:尽量增大,批量内负样本越多,对比学习效果越好,内存不足时可使用梯度累积(
gradient_accumulation_steps)。
- 评估方式:不要用分类指标,改用嵌入相似性相关指标,比如余弦相似度、Recall@k、MRR等,在验证集上测试查询样本与候选样本的匹配准确率。
- 模型保存与复用:训练完成后,可提取
model.backbone作为微调后的Sentence Transformer模型,直接用于嵌入生成任务。
内容的提问来源于stack exchange,提问作者Rudranath Palit
相关产品推荐
相关产品推荐

