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

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记录训练损失,完全不符合损失累积、计算平均的基本逻辑,无法正确监控训练状态。
修正建议
  1. 改用双塔模型标准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]
    
  2. 匹配对应的损失函数
    使用带温度的交叉熵损失(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()
    
  3. 修复损失日志逻辑

    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
    
  4. 保留原始数据类型
    直接用fields的原始整数类型输入模型,不要做多余的类型转换。

内容的提问来源于stack exchange,提问作者Junyeong Choi

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.09 11:43:10