PyTorch批内负采样实现及双塔模型训练性能优化求助
问题描述
我用隐式数据集训练双塔推荐模型,想在训练阶段通过in-batch负采样优化训练效果。自己写的训练代码性能极差,根本无法完成训练;之前采用随机负采样的训练效果也不理想,求帮忙检查训练逻辑。
训练代码
def train_with_in_batch_negative_sampling( model, num_epoch, optimizer, data_loader, criterion, device, log_interval=10, num_neg_samples=5, ): model.train() total_loss = 0 train_loss = 987654321 tk0 = tqdm.tqdm(data_loader, smoothing=0, mininterval=1.0) for i, (fields, target) in enumerate(tk0): fields, target = fields.to(device), target.to(device) new_fields = [] for idx, row in enumerate(fields): if idx==0: item_tensor = fields[idx+1:,1] else: item_tensor = torch.cat([fields[:idx,1], fields[idx+1:,1]],dim=0) new_fields.append(torch.Tensor([row[0], row[1], torch.Tensor(1)]).view(1, -1).to(device)) new_fields.append( torch.cat( [ row[0:1].repeat(511).view(-1, 1), item_tensor.view(-1, 1), torch.zeros(511, 1).to(device), ], dim=1, ) ) merged = torch.cat(new_fields,dim=0).to(torch.int) y = model(merged[:,:2]) loss = criterion(y, merged[:,2]) model.zero_grad() loss.backward() optimizer.step() total_loss += loss.item() train_loss = min(train_loss, total_loss / log_interval) if (i + 1) % log_interval == 0: tk0.set_postfix(loss=total_loss / log_interval, epoch_num=num_epoch) total_loss = 0 return train_loss
核心问题分析
- in-batch负采样逻辑完全错误:
你硬编码了511个负样本数量,但实际负样本数应该是batch_size-1(当前batch内除自身正样本外的所有item),一旦batch_size不等于512,就会出现维度不匹配,直接导致无效计算或报错。同时手动构造样本的方式会让batch规模暴增为原有的512倍,计算量爆炸,这是性能极差的核心原因。 - 数据类型处理错误:
构造样本时先转成浮点型torch.Tensor,最后又强制转成torch.int,会导致用户ID、item ID等离散特征被截断,模型输入完全失真。 - 损失计算逻辑不匹配双塔模型:
双塔模型的in-batch负采样应该基于用户/物品embedding的相似度矩阵计算损失(比如InfoNCE、带温度的交叉熵),你直接用二元标签喂给普通损失函数,会因为正负样本数量失衡、输出与标签维度不匹配导致损失计算无效。 - 损失日志与更新逻辑混乱:
每次迭代后直接清零total_loss,且用min记录训练损失,完全不符合损失累积、计算平均的基本逻辑,无法正确监控训练状态。
修正建议
改用双塔模型标准in-batch负采样实现
不需要手动构造负样本,直接利用当前batch的用户/物品embedding计算相似度矩阵,对角线为正样本,其余为负样本:# 假设模型分用户塔和物品塔 user_emb = model.user_tower(fields[:, 0]) # shape: [batch_size, emb_dim] item_emb = model.item_tower(fields[:, 1]) # shape: [batch_size, emb_dim] sim_matrix = torch.matmul(user_emb, item_emb.T) # shape: [batch_size, batch_size]匹配对应的损失函数
使用带温度的交叉熵损失(InfoNCE的简化版),把每个用户的正样本作为目标:temperature = 0.1 sim_matrix = sim_matrix / temperature labels = torch.arange(fields.size(0), device=device) # 正样本是对角线位置 loss = criterion(sim_matrix, labels) # criterion用torch.nn.CrossEntropyLoss()修复损失日志逻辑
total_loss += loss.item() if (i + 1) % log_interval == 0: avg_loss = total_loss / log_interval tk0.set_postfix(loss=avg_loss, epoch_num=num_epoch) train_loss = min(train_loss, avg_loss) total_loss = 0保留原始数据类型
直接用fields的原始整数类型输入模型,不要做多余的类型转换。
内容的提问来源于stack exchange,提问作者Junyeong Choi
相关产品推荐
相关产品推荐

