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()
- 方案一:将one-hot标签转为类别索引后使用
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
相关产品推荐
相关产品推荐

