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

PyTorch二分类模型训练Loss持续下降问题排查求助

二分类模型训练异常排查:训练Loss持续下降、测试Loss波动

我用PyTorch训练二分类模型,输入向量长度561(341维是one-hot编码,其余是0-1区间特征),输出为[0,1]或[1,0]。目前遇到的问题是:训练Loss一直持续下降,哪怕训到200个epoch也没有收敛迹象;同时测试Loss在下降和上升之间波动。怀疑是Loss计算或者训练流程有问题。

试过换不同模型结构(LSTM、CNN)、更换损失函数,问题依然存在。以下是相关代码和训练结果:

模型代码

class MyRegression(nn.Module):
    def __init__(self, input_dim, output_dim):
        super(MyRegression, self).__init__()
        # One layer
        self.linear1 = nn.Linear(input_dim, 128)
        self.linear2 = nn.Linear(128, output_dim)
    def forward(self, x):
        return self.linear2(self.linear1(x))

训练函数

def run_gradient_descent(model, data_train, data_val, batch_size, learning_rate, weight_decay=0, num_epochs=20):
    
    model = model.to(device)
    criterion = nn.CrossEntropyLoss()
    #criterion = nn.MSELoss()
    optimizer = optim.SGD(model.parameters(), lr=learning_rate, weight_decay=weight_decay)
    iters, losses, train_losses, test_losses = [], [], [], []
    iters_sub, train_acc, val_acc = [], [] ,[]
    print(batch_size)
    
    # weight sampler
    class0, class1 =labels_count(data_train)
    dataset_counts = [class0, class1]
    print(dataset_counts)
    num_samples = sum(dataset_counts)
    labels = [tag for _, tag in data_train]
    #max_value = max(input_list)
    #index = input_list.index(max_value)
    class_weights = [1./dataset_counts[i] for i in range(len(dataset_counts))]
    labels_indics = [i.index(max(i)) for i in labels ]
    weights = [class_weights[i] for i in labels_indics] # labels.max(1, keepdim=True)[1]
    weights = numpy.array(weights)
    samples_weight = torch.from_numpy(weights)
    samples_weigth = samples_weight.double()
    sampler = torch.utils.data.sampler.WeightedRandomSampler(samples_weight, int(num_samples), replacement=True)
    
    
    train_loader = torch.utils.data.DataLoader(
        data_train,
        batch_size=batch_size,
        shuffle=False,
        sampler = sampler,
        collate_fn=lambda d: ([x[0] for x in d], [x[1] for x in d]),
        num_workers=os.cpu_count()//2
    )

    # training
    n = 0 # the number of iterations
    for epoch in tqdm(range(num_epochs), desc="epoch"):
        correct = 0
        total = 0
        for xs, ts in tqdm(train_loader, desc="train"):
            xs = torch.FloatTensor(xs).to(device)
            ts = torch.FloatTensor(ts).to(device)
            # print("batch index {}, 0/1: {}/{}".format(n,ts.tolist().count([1,0]),ts.tolist().count([0,1])))

            # if len(ts) != batch_size:
            #     print("ops")
            #     continue
            model.train()
            
            zs = model(xs)
            zs = zs.to(device)
            loss = criterion(zs, ts)
            loss.backward()
            optimizer.step()
            optimizer.zero_grad()

            iters.append(n)
            loss.detach().cpu()
            
            losses.append(float(loss)/len(ts)) # compute *average* loss
            
            pred = zs.max(1, keepdim=True)[1] # get the index of the max logit
            target = ts.max(1, keepdim=True)[1]
            correct += pred.eq(target).sum().item()
            total += int(ts.shape[0])
            acc = correct / total

            if (n % len(train_loader) == 0) and n>0 and epoch%2==0:
                
                test_acc, test_loss = get_accuracy(model, data_val)
            
                iters_sub.append(n)
                train_acc.append(acc)
                val_acc.append(test_acc)
        
                train_losses.append(sum(losses)/len(losses))
                test_losses.append(test_loss)
                
                print("Epoch", epoch, "train_acc", acc)
                print("Epoch", epoch, "test_acc", test_acc)
                print("Epoch", epoch, "train_loss", sum(losses)/len(losses))
                print("Epoch", epoch, "test_loss", test_loss)

             # increment the iteration number
            n += 1
        torch.save(model.state_dict(), f"{MODEL_NAME}/checkpoint_epoch{epoch}.pt")


    # plotting
    plt.title("Training Curve (batch_size={}, lr={})".format(batch_size, learning_rate))
    plt.plot(iters_sub, train_losses, label="Train")
    plt.plot(iters_sub, test_losses, label="Test")
    plt.legend(loc='best')
    plt.xlabel("Iterations")
    plt.ylabel("Loss")
    
    plt.savefig(f"{MODEL_NAME}/training_test_loss.png")
    
    # plt.show()
    plt.clf()
    plt.title("Training Curve (batch_size={}, lr={})".format(batch_size, learning_rate))
    plt.plot(iters_sub, train_acc, label="Train")
    plt.plot(iters_sub, val_acc, label="Test")
    plt.xlabel("Iterations")
    plt.ylabel("Accuracy")
    plt.legend(loc='best')
    plt.savefig(f"{MODEL_NAME}/training_acc.png")
    #plt.show()
    return model

主函数

model = MyRegression(374, 2)
run_gradient_descent(
    model,
    training_set,
    test_set,
    batch_size= 64,  
    learning_rate=1e-2,
    num_epochs=200
)

部分训练结果

Epoch 2 train_acc 0.578125
Epoch 2 test_acc 0.7346171218510883
Epoch 2 train_loss 0.003494985813946325
Epoch 2 test_loss 0.00318981208993754
Epoch 4 train_acc 0.671875
Epoch 4 test_acc 0.7021743310868525
Epoch 4 train_loss 0.0034714722261212196
Epoch 4 test_loss 0.0033061892530283398
Epoch 6 train_acc 0.75
Epoch 6 test_acc 0.7614966302787455
Epoch 6 train_loss 0.003462064279302097
Epoch 6 test_loss 0.003087314312623757
Epoch 8 train_acc 0.625
Epoch 8 test_acc 0.7343577405202831
Epoch 8 train_loss 0.0034565126970269753
Epoch 8 test_loss 0.0032059013449951632
Epoch 10 train_acc 0.578125
Epoch 10 test_acc 0.7587194612023667
Epoch 10 train_loss 0.0034528369772701857
Epoch 10 test_loss 0.003112017690331294
Epoch 12 train_acc 0.65625
Epoch 12 test_acc 0.7097187501397528
Epoch 12 train_loss 0.003450584381555143
Epoch 12 test_loss 0.003285413007535127
Epoch 14 train_acc 0.578125
Epoch 14 test_acc 0.7509648538296759
Epoch 14 train_loss 0.0034486886994226553
Epoch 14 test_loss 0.003145160475069196
Epoch 16 train_acc 0.625
Epoch 16 test_acc 0.7629612403794123
Epoch 16 train_loss 0.0034474354597715125
Epoch 16 test_loss 0.003106232365138448
Epoch 18 train_acc 0.703125
Epoch 18 test_acc 0.7527134417666552
Epoch 18 train_loss 0.0034464063646294537
Epoch 18 test_loss 0.0031368749897371824

核心问题排查与修复建议

1. 损失函数与标签格式完全不匹配

CrossEntropyLoss要求目标标签是类别索引(如0或1),但你传入的是one-hot编码的[0,1]/[1,0],这会直接导致Loss计算逻辑错误,是核心问题。

  • 修复方案:
    • 方案一:将one-hot标签转为类别索引后使用CrossEntropyLoss:
      # 训练循环内转换标签格式
      ts = torch.argmax(ts, dim=1).to(device)
      criterion = nn.CrossEntropyLoss()
      
    • 方案二:保留one-hot标签,改用BCEWithLogitsLoss(模型输出无需加激活,损失内部会处理sigmoid):
      criterion = nn.BCEWithLogitsLoss()
      

2. 输入维度严重不匹配

主函数中模型初始化用了MyRegression(374, 2),但实际输入向量长度是561,维度不匹配会导致模型学习完全偏离预期。

  • 修复:将模型初始化改为MyRegression(561, 2)

3. Loss记录逻辑错误

  • CrossEntropyLoss默认已经是batch平均损失(reduction='mean'),你额外除以len(ts)会导致Loss数值被过度缩小(训练结果中Loss仅0.003左右,完全不符合正常范围)。
  • loss.detach().cpu()这行代码没有赋值,完全无效,应该提取数值后再记录。
  • 修复:
    loss_val = loss.detach().cpu().item()
    losses.append(loss_val)  # 无需再除以batch size
    

4. 加权采样存在笔误

代码中samples_weigth = samples_weight.double()写错了变量名(weigth→weight),导致采样器未使用double类型的权重,可能影响类别平衡效果。

  • 修复:
    samples_weight = torch.from_numpy(weights).double()
    sampler = torch.utils.data.sampler.WeightedRandomSampler(samples_weight, int(num_samples), replacement=True)
    

5. 模型缺少中间激活层

当前模型是两层线性层直接堆叠,没有非线性激活,相当于单个线性层,无法学习复杂特征。

  • 修复:在linear1后添加激活函数:
    def forward(self, x):
        x = self.linear1(x)
        x = torch.relu(x)  # 添加ReLU激活
        return self.linear2(x)
    

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.29 04:44:59