如何用PyTorch向量化实现有监督SimCSE损失计算?
有监督SimCSE损失的向量化优化实现(PyTorch)
完全可以通过广播操作与矩阵乘法实现有监督SimCSE损失的向量化优化,彻底替代朴素实现中的循环逻辑,大幅提升计算效率,尤其适合大批次训练场景。
朴素实现回顾(供对比)
通常朴素实现会通过逐样本循环计算损失,示例代码如下:
import torch import torch.nn.functional as F def naive_supervised_simcse_loss(embeddings, labels, temperature=0.05): batch_size = embeddings.size(0) loss = 0.0 embeddings = F.normalize(embeddings, p=2, dim=1) for i in range(batch_size): # 筛选当前样本的正样本(同标签且非自身) pos_mask = (labels == labels[i]) & (torch.arange(batch_size) != i) # 计算当前样本与所有样本的相似度 sims = torch.matmul(embeddings[i].unsqueeze(0), embeddings.T).squeeze() / temperature # 累加正样本相似度,计算分母的logsumexp pos_sim_sum = sims[pos_mask].sum() log_denominator = torch.logsumexp(sims, dim=0) loss += (-pos_sim_sum + log_denominator) return loss / batch_size
这种实现的问题在于循环无法利用GPU并行计算能力,批次越大,计算效率越低。
向量化优化实现
基于广播与矩阵乘法的优化思路,核心是一次性计算全批次的相似度矩阵,再通过掩码快速筛选正样本:
import torch import torch.nn.functional as F def vectorized_supervised_simcse_loss(embeddings, labels, temperature=0.05): batch_size = embeddings.size(0) # 归一化嵌入向量,确保相似度计算为余弦相似度 embeddings = F.normalize(embeddings, p=2, dim=1) # 一次性计算全批次相似度矩阵,形状为 (batch_size, batch_size) sim_matrix = torch.matmul(embeddings, embeddings.T) / temperature # 构建正样本掩码:同标签且排除自身 # 通过广播实现标签的两两对比,再结合单位矩阵排除自身样本 pos_mask = (labels.unsqueeze(0) == labels.unsqueeze(1)) & ~torch.eye(batch_size, device=embeddings.device, dtype=torch.bool) # 计算每个样本的正样本相似度总和 pos_sim_sum = (sim_matrix * pos_mask).sum(dim=1) # 计算每个样本的分母项:所有样本相似度的logsumexp log_denominator = torch.logsumexp(sim_matrix, dim=1) # 计算平均损失 loss = (-pos_sim_sum + log_denominator).mean() return loss
优化说明
- 全批次相似度计算:通过一次矩阵乘法得到所有样本间的相似度,避免循环中重复计算
- 广播构建掩码:利用
labels.unsqueeze(0)与labels.unsqueeze(1)的广播特性,快速生成两两标签对比矩阵,再结合取反的单位矩阵排除自身 - 并行化求和与logsumexp:对相似度矩阵的行直接执行求和与logsumexp操作,充分利用GPU并行计算能力
正确性验证
可以通过以下代码验证两种实现的输出一致性:
# 生成测试输入 embeddings = torch.randn(32, 768) labels = torch.randint(0, 5, (32,)) # 计算两种实现的损失 loss_naive = naive_supervised_simcse_loss(embeddings, labels) loss_vectorized = vectorized_supervised_simcse_loss(embeddings, labels) print(f"朴素实现损失: {loss_naive.item():.6f}") print(f"向量化实现损失: {loss_vectorized.item():.6f}") print(f"损失差值: {torch.abs(loss_naive - loss_vectorized).item():.10f}")
运行后会看到两者损失值几乎完全一致,仅存在微小浮点数精度差异。
内容的提问来源于stack exchange,提问作者Gonzo
相关产品推荐
相关产品推荐

