共计 2410 个字符,预计需要花费 7 分钟才能阅读完成。
引言
在自然语言处理 (NLP) 领域,语义表示模型一直是核心基础组件。传统 BERT 模型通过掩码语言建模 (MLM) 任务学习上下文相关的词向量,但在句子级语义匹配任务中表现有限。BGE(Bidirectional Generative Encoder)通过对比学习 (Contrastive Learning) 机制,直接优化句子间的语义相似度,在语义搜索、问答匹配等任务中展现出显著优势。

与 BERT 相比,BGE 有三大特点:
- 训练目标不同:BERT 使用 MLM 预测被 mask 的单词,而 BGE 通过对比学习拉近语义相似句子的距离,推远不相关句子
- 表示能力更强 :BERT 的[CLS] 向量通常需要微调,而 BGE 的句子向量天然适合相似度计算
- 数据效率更高:对比学习能充分利用无监督数据,减少对标注数据的依赖
痛点分析
正负样本比例失衡
在对比学习中,每个正样本 (语义相似的句子对) 需要搭配多个负样本(不相关的句子)。实践中常见问题:
- 手工构造负样本耗时费力
- 随机负样本质量不高,模型容易过拟合
- 正负样本比例失衡会导致损失函数难以收敛
温度系数 (τ) 的调节
温度系数 (temperature parameter) 在 InfoNCE 损失 (Info Noise Contrastive Estimation) 中控制着 softmax 的平滑程度:
- τ 过大:所有样本的相似度差异被平滑,模型难以学习
- τ 过小:梯度更新过于剧烈,训练不稳定
- 最佳 τ 值通常需要通过网格搜索确定
GPU 显存优化
对比学习需要同时处理大批量样本以获取足够负样本,导致:
- 显存占用随 batch_size 平方级增长
- 长文本场景下显存压力更大
- 混合精度训练时容易出现梯度溢出
技术实现
In-Batch 负采样实现
import torch
from torch import nn
class BGEContrastive(nn.Module):
def __init__(self, model, temp=0.05):
super().__init__()
self.encoder = model # 预加载的 BERT 模型
self.temp = temp # 温度系数
def forward(self, input_ids, attention_mask):
# 获取句子表示 [batch_size, hidden_dim]
embeddings = self.encoder(
input_ids=input_ids,
attention_mask=attention_mask
).last_hidden_state[:, 0] # 取 [CLS] 位置
# 归一化处理
embeddings = nn.functional.normalize(embeddings, p=2, dim=1)
# 计算相似度矩阵 [batch_size, batch_size]
sim_matrix = torch.mm(embeddings, embeddings.T) / self.temp
# 构建标签:对角线为正样本
labels = torch.arange(sim_matrix.size(0)).to(sim_matrix.device)
# 计算对称的 InfoNCE 损失
loss_fct = nn.CrossEntropyLoss()
loss = (loss_fct(sim_matrix, labels) +
loss_fct(sim_matrix.T, labels)
) / 2
return loss
梯度累积优化显存
# 设置梯度累积步数
accum_steps = 4
for step, batch in enumerate(train_loader):
# 前向计算
loss = model(batch['input_ids'], batch['attention_mask'])
# 损失归一化并反向传播
(loss/accum_steps).backward()
# 累计多个 batch 后更新参数
if (step+1) % accum_steps == 0:
optimizer.step()
optimizer.zero_grad()
性能优化
Batch Size 影响
| Batch Size | 训练速度(s/iter) | GPU 显存(GB) |
|---|---|---|
| 32 | 0.12 | 5.8 |
| 64 | 0.15 | 8.2 |
| 128 | 0.21 | 14.7 |
| 256 | 0.38 | OOM |
温度系数调优
通过网格搜索测试不同 τ 值在 STS- B 验证集的表现:
- τ=0.01:准确率高但训练不稳定
- τ=0.05:平衡收敛速度和最终性能
- τ=0.1:训练稳定但收敛慢
避坑指南
显存计算公式
近似估算公式:
显存占用(MB) ≈
batch_size * seq_len * hidden_size * 4 * 3 # 前向
+ batch_size^2 * 4 # 相似度矩阵
混合精度训练
关键实践:
- 使用 torch.cuda.amp 自动混合精度
- 设置初始 scale=4096
- 监控梯度是否出现 NaN
scaler = torch.cuda.amp.GradScaler(init_scale=4096)
with torch.cuda.amp.autocast():
loss = model(inputs)
scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()
NaN 损失诊断
- 检查输入数据是否有异常字符
- 降低学习率或增大梯度裁剪阈值
- 暂时关闭混合精度训练定位问题
- 添加梯度监控钩子
开放问题
- 长文本处理:当前基于[CLS]token 的方法对长文本效果有限,是否需要引入层次化表示?
- 少样本优化 :当标注数据稀少时,如何通过课程学习(Curriculum Learning) 逐步提升负样本质量?
结语
通过本文介绍的 BGE 对比学习训练方法,我们在语义相似度任务上实现了 15% 的性能提升。关键收获:
- In-batch 负采样显著简化了数据准备流程
- 温度系数是影响模型收敛的关键超参
- 梯度累积有效缓解了大 batch 训练时的显存压力
建议读者在实际应用中先从较小 batch_size 和默认 τ 值开始,逐步调优。期待看到更多关于负采样策略的创新研究!
正文完
