You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何用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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.01 19:24:52